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
|
|
|
*/
|
2017-11-22 21:16:10 +00:00
|
|
|
public function getTargetClass($target): int
|
2017-05-17 22:07:14 +00:00
|
|
|
{
|
2018-02-16 06:25:24 +00:00
|
|
|
if (!in_array($target, $this->classes, true)) {
|
2018-03-03 15:03:53 +00:00
|
|
|
throw new InvalidArgumentException(
|
|
|
|
sprintf('Target with value "%s" is not part of the accepted classes', $target)
|
|
|
|
);
|
2017-05-17 22:07:14 +00:00
|
|
|
}
|
2017-08-17 06:50:37 +00:00
|
|
|
|
2018-02-16 06:25:24 +00:00
|
|
|
return array_search($target, $this->classes, true);
|
2017-05-17 22:07:14 +00:00
|
|
|
}
|
|
|
|
|
|
|
|
/**
|
|
|
|
* @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;
|
|
|
|
}
|
|
|
|
}
|
2017-08-17 06:50:37 +00:00
|
|
|
|
2017-05-17 22:07:14 +00:00
|
|
|
return $this->classes[$predictedClass];
|
|
|
|
}
|
|
|
|
|
|
|
|
/**
|
|
|
|
* @param mixed $target
|
|
|
|
*/
|
2017-11-14 20:21:23 +00:00
|
|
|
protected function trainSample(array $sample, $target): void
|
2017-05-17 22:07:14 +00:00
|
|
|
{
|
|
|
|
|
|
|
|
// Feed-forward.
|
|
|
|
$this->setInput($sample)->getOutput();
|
|
|
|
|
|
|
|
// Back-propagate.
|
|
|
|
$this->backpropagation->backpropagate($this->getLayers(), $this->getTargetClass($target));
|
|
|
|
}
|
|
|
|
}
|