2017-05-17 22:07:14 +00:00
|
|
|
<?php
|
|
|
|
|
|
|
|
declare(strict_types=1);
|
|
|
|
|
|
|
|
namespace Phpml\Classification;
|
|
|
|
|
|
|
|
use Phpml\Exception\InvalidArgumentException;
|
|
|
|
use Phpml\NeuralNetwork\Network\MultilayerPerceptron;
|
|
|
|
|
|
|
|
class MLPClassifier extends MultilayerPerceptron implements Classifier
|
|
|
|
{
|
|
|
|
|
|
|
|
/**
|
2017-07-26 06:22:12 +00:00
|
|
|
* @param mixed $target
|
|
|
|
*
|
|
|
|
* @throws InvalidArgumentException
|
|
|
|
*
|
2017-05-17 22:07:14 +00:00
|
|
|
* @return int
|
|
|
|
*/
|
|
|
|
public function getTargetClass($target): int
|
|
|
|
{
|
|
|
|
if (!in_array($target, $this->classes)) {
|
|
|
|
throw InvalidArgumentException::invalidTarget($target);
|
|
|
|
}
|
|
|
|
return array_search($target, $this->classes);
|
|
|
|
}
|
|
|
|
|
|
|
|
/**
|
|
|
|
* @param array $sample
|
|
|
|
*
|
|
|
|
* @return mixed
|
|
|
|
*/
|
|
|
|
protected function predictSample(array $sample)
|
|
|
|
{
|
|
|
|
$output = $this->setInput($sample)->getOutput();
|
|
|
|
|
|
|
|
$predictedClass = null;
|
|
|
|
$max = 0;
|
|
|
|
foreach ($output as $class => $value) {
|
|
|
|
if ($value > $max) {
|
|
|
|
$predictedClass = $class;
|
|
|
|
$max = $value;
|
|
|
|
}
|
|
|
|
}
|
|
|
|
return $this->classes[$predictedClass];
|
|
|
|
}
|
|
|
|
|
|
|
|
/**
|
|
|
|
* @param array $sample
|
|
|
|
* @param mixed $target
|
|
|
|
*/
|
|
|
|
protected function trainSample(array $sample, $target)
|
|
|
|
{
|
|
|
|
|
|
|
|
// Feed-forward.
|
|
|
|
$this->setInput($sample)->getOutput();
|
|
|
|
|
|
|
|
// Back-propagate.
|
|
|
|
$this->backpropagation->backpropagate($this->getLayers(), $this->getTargetClass($target));
|
|
|
|
}
|
|
|
|
}
|