Cancel Training
This commit is contained in:
@@ -0,0 +1,9 @@
|
||||
<?php
|
||||
|
||||
namespace App\Exceptions;
|
||||
|
||||
use RuntimeException;
|
||||
|
||||
class TrainingCancelledException extends RuntimeException
|
||||
{
|
||||
}
|
||||
@@ -3,6 +3,7 @@
|
||||
namespace App\Http\Controllers;
|
||||
|
||||
use App\Events\PerceptronInitialization;
|
||||
use App\Exceptions\TrainingCancelledException;
|
||||
use App\Http\Requests\RunPerceptronRequest;
|
||||
use App\Models\NetworksTraining\ADALINEPerceptronTraining;
|
||||
use App\Models\NetworksTraining\GradientDescentPerceptronTraining;
|
||||
@@ -18,9 +19,25 @@ use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use App\Services\SynapticWeightsProvider\RandomSynapticWeights;
|
||||
use App\Services\SynapticWeightsProvider\ZeroSynapticWeights;
|
||||
use Illuminate\Http\Request;
|
||||
use Illuminate\Support\Facades\Cache;
|
||||
|
||||
class PerceptronController extends Controller
|
||||
{
|
||||
private function cancellationKey(string $trainingId): string
|
||||
{
|
||||
return "perceptron-training-cancelled:{$trainingId}";
|
||||
}
|
||||
|
||||
public function cancel(Request $request)
|
||||
{
|
||||
$trainingId = $request->validate([
|
||||
'training_id' => ['required', 'string', 'max:100'],
|
||||
])['training_id'];
|
||||
|
||||
Cache::put($this->cancellationKey($trainingId), true, now()->addHour());
|
||||
|
||||
return response()->noContent();
|
||||
}
|
||||
/**
|
||||
* Display the specified resource.
|
||||
*/
|
||||
@@ -220,6 +237,8 @@ class PerceptronController extends Controller
|
||||
$sessionId = $request->input('session_id', session()->getId());
|
||||
$trainingId = $request->input('training_id');
|
||||
|
||||
Cache::forget($this->cancellationKey($trainingId));
|
||||
|
||||
// Zero initialization prevents hidden layers from receiving a gradient.
|
||||
if ($perceptronType === 'multilayer' && $weightInitMethod === 'zeros') {
|
||||
$synapticWeightsProvider = new RandomSynapticWeights;
|
||||
@@ -236,17 +255,22 @@ class PerceptronController extends Controller
|
||||
$datasetReader = $this->getDataSetReader($dataSet);
|
||||
|
||||
$networkTraining = match ($perceptronType) {
|
||||
'simple' => new SimpleBinaryPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId),
|
||||
'gradientdescent' => new GradientDescentPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError),
|
||||
'adaline' => new ADALINEPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError),
|
||||
'monolayer' => new MonoLayerPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError),
|
||||
'multilayer' => new MultiLayerPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $hiddenLayers, $hiddenLayersNeurons, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError),
|
||||
'simple' => new SimpleBinaryPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, fn (): bool => connection_aborted() || Cache::has($this->cancellationKey($trainingId))),
|
||||
'gradientdescent' => new GradientDescentPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError, fn (): bool => connection_aborted() || Cache::has($this->cancellationKey($trainingId))),
|
||||
'adaline' => new ADALINEPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError, fn (): bool => connection_aborted() || Cache::has($this->cancellationKey($trainingId))),
|
||||
'monolayer' => new MonoLayerPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError, fn (): bool => connection_aborted() || Cache::has($this->cancellationKey($trainingId))),
|
||||
'multilayer' => new MultiLayerPerceptronTraining($datasetReader, $learningRate, $maxEpochs, $hiddenLayers, $hiddenLayersNeurons, $synapticWeightsProvider, $iterationEventBuffer, $sessionId, $trainingId, $minError, fn (): bool => connection_aborted() || Cache::has($this->cancellationKey($trainingId))),
|
||||
default => null,
|
||||
};
|
||||
|
||||
event(new PerceptronInitialization($datasetReader->lines, $networkTraining->activationFunction, $sessionId, $trainingId));
|
||||
|
||||
$networkTraining->start();
|
||||
try {
|
||||
$networkTraining->start();
|
||||
} catch (TrainingCancelledException) {
|
||||
$networkTraining->cancel();
|
||||
Cache::forget($this->cancellationKey($trainingId));
|
||||
}
|
||||
|
||||
return back()->with('success', [
|
||||
'message' => 'Training completed',
|
||||
|
||||
@@ -8,6 +8,7 @@ use App\Models\Perceptrons\Perceptron;
|
||||
use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use Closure;
|
||||
|
||||
class ADALINEPerceptronTraining extends NetworkTraining
|
||||
{
|
||||
@@ -26,8 +27,9 @@ class ADALINEPerceptronTraining extends NetworkTraining
|
||||
string $sessionId,
|
||||
string $trainingId,
|
||||
private float $minError,
|
||||
?Closure $isCancelled = null,
|
||||
) {
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId);
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId, $isCancelled);
|
||||
$this->perceptron = new GradientDescentPerceptron($synapticWeightsProvider->generate($datasetReader->getInputSize()));
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ use App\Models\Perceptrons\Perceptron;
|
||||
use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use Closure;
|
||||
|
||||
class GradientDescentPerceptronTraining extends NetworkTraining
|
||||
{
|
||||
@@ -26,8 +27,9 @@ class GradientDescentPerceptronTraining extends NetworkTraining
|
||||
string $sessionId,
|
||||
string $trainingId,
|
||||
private float $minError,
|
||||
?Closure $isCancelled = null,
|
||||
) {
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId);
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId, $isCancelled);
|
||||
$this->perceptron = new GradientDescentPerceptron($synapticWeightsProvider->generate($datasetReader->getInputSize()));
|
||||
}
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use App\Services\SynapticWeightsProvider\SimpleNetworkWeightsProvider;
|
||||
use Closure;
|
||||
use Illuminate\Support\Arr;
|
||||
|
||||
class MonoLayerPerceptronTraining extends NetworkTraining
|
||||
@@ -35,8 +36,9 @@ class MonoLayerPerceptronTraining extends NetworkTraining
|
||||
string $sessionId,
|
||||
string $trainingId,
|
||||
private float $minError,
|
||||
?Closure $isCancelled = null,
|
||||
) {
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId);
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId, $isCancelled);
|
||||
$this->isRegression = $datasetReader->getInputSize() === 1;
|
||||
$networkWeightsProvider = new SimpleNetworkWeightsProvider($synapticWeightsProvider);
|
||||
$this->network = new NetworkPerceptron(
|
||||
|
||||
@@ -11,6 +11,7 @@ use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use App\Services\SynapticWeightsProvider\SimpleNetworkWeightsProvider;
|
||||
use Closure;
|
||||
use Illuminate\Support\Arr;
|
||||
|
||||
class MultiLayerPerceptronTraining extends NetworkTraining
|
||||
@@ -37,8 +38,9 @@ class MultiLayerPerceptronTraining extends NetworkTraining
|
||||
string $sessionId,
|
||||
string $trainingId,
|
||||
private float $minError,
|
||||
?Closure $isCancelled = null,
|
||||
) {
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId);
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId, $isCancelled);
|
||||
$this->labels = $datasetReader->getLabels();
|
||||
$this->isRegression = $datasetReader->getOutputSize() === 1
|
||||
|| ($datasetReader->getOutputSize() > 2
|
||||
|
||||
@@ -3,9 +3,11 @@
|
||||
namespace App\Models\NetworksTraining;
|
||||
|
||||
use App\Events\PerceptronTrainingEnded;
|
||||
use App\Exceptions\TrainingCancelledException;
|
||||
use App\Models\ActivationsFunctions;
|
||||
use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use Closure;
|
||||
|
||||
abstract class NetworkTraining
|
||||
{
|
||||
@@ -24,6 +26,7 @@ abstract class NetworkTraining
|
||||
protected IPerceptronIterationEventBuffer $iterationEventBuffer,
|
||||
protected string $sessionId,
|
||||
protected string $trainingId,
|
||||
protected ?Closure $isCancelled = null,
|
||||
) {}
|
||||
|
||||
abstract public function start(): void;
|
||||
@@ -50,9 +53,18 @@ abstract class NetworkTraining
|
||||
|
||||
protected function addIterationToBuffer(float $error, array $synapticWeights)
|
||||
{
|
||||
if ($this->isCancelled !== null && ($this->isCancelled)()) {
|
||||
throw new TrainingCancelledException;
|
||||
}
|
||||
|
||||
$this->iterationEventBuffer->addIteration($this->epoch, $this->datasetReader->getLastReadLineIndex(), $error, $synapticWeights);
|
||||
}
|
||||
|
||||
public function cancel(): void
|
||||
{
|
||||
$this->broadcastTrainingEnded('Entraînement annulé');
|
||||
}
|
||||
|
||||
public function getEpoch(): int
|
||||
{
|
||||
return $this->epoch;
|
||||
|
||||
@@ -9,6 +9,7 @@ use App\Models\Perceptrons\SimpleBinaryPerceptron;
|
||||
use App\Services\DatasetReader\IDataSetReader;
|
||||
use App\Services\IterationEventBuffer\IPerceptronIterationEventBuffer;
|
||||
use App\Services\SynapticWeightsProvider\ISynapticWeightsProvider;
|
||||
use Closure;
|
||||
|
||||
class SimpleBinaryPerceptronTraining extends NetworkTraining
|
||||
{
|
||||
@@ -28,8 +29,9 @@ class SimpleBinaryPerceptronTraining extends NetworkTraining
|
||||
IPerceptronIterationEventBuffer $iterationEventBuffer,
|
||||
string $sessionId,
|
||||
string $trainingId,
|
||||
?Closure $isCancelled = null,
|
||||
) {
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId);
|
||||
parent::__construct($datasetReader, $maxEpochs, $iterationEventBuffer, $sessionId, $trainingId, $isCancelled);
|
||||
$this->perceptron = new SimpleBinaryPerceptron($synapticWeightsProvider->generate($datasetReader->getInputSize()));
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user