feat: models caching

This commit is contained in:
2026-08-10 15:09:04 +03:00
parent 67718629ae
commit ec054d88c4
2 changed files with 179 additions and 114 deletions
+1
View File
@@ -3,6 +3,7 @@
## ##
## Get latest from `dotnet new gitignore` ## Get latest from `dotnet new gitignore`
.models_cache
.zed .zed
*.png *.png
cache cache
+178 -114
View File
@@ -141,69 +141,102 @@ public partial class AnalyzerService
var sampleRecords = _context.SampleRecords var sampleRecords = _context.SampleRecords
.Where(r => r.BrandId == brand.Id); .Where(r => r.BrandId == brand.Id);
_logger.LogInformation("Найдено калибровочных записей проб для марки: {}", sampleRecords.Count()); var dataSb = new System.Text.StringBuilder();
if (sampleRecords.Count() == 0) dataSb.AppendLine(brand.Id.ToString());
throw new InvalidDataException($"Не найдено калибровочных записей для марки: \"{brand.BrandName} - {brand.MixtureName}\""); dataSb.AppendLine(brand.EditDateTime.ToString());
cancellationToken.ThrowIfCancellationRequested(); foreach (var record in sampleRecords.OrderBy(r => r.Id))
if (separator.Features.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных признаков для сепаратора: \"{separator.Type}\"");
if (separator.Features.Count() % FEATURES_LENGTH != 0)
throw new InvalidDataException($"Неожидаемое количество калибровочных признаков для сепаратора: \"{separator.Type}\"");
var separatorDataSet = separator.Features.Chunk(FEATURES_LENGTH)
.Select(v => new ClusterizationData
{
Label = "separator",
Features = v.ToArray()
})
.Shuffle();
_logger.LogInformation("Загружено векторов для сепаратора: {}", separatorDataSet.Count());
if (separatorDataSet.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных данных для сепаратора: \"{separator.Type}\"");
cancellationToken.ThrowIfCancellationRequested();
foreach (var sampleRecord in sampleRecords)
{ {
if (sampleRecord.Features.Count() == 0) dataSb.AppendLine(record.Id.ToString());
_logger.LogError("Не найдено калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName); dataSb.AppendLine(record.EditDateTime.ToString());
if (sampleRecord.Features.Count() % FEATURES_LENGTH != 0)
_logger.LogError("Неожидаемое количество калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
} }
dataSb.AppendLine(separator.Id.ToString());
dataSb.AppendLine(separator.EditDateTime.ToString());
var dataHashCode = Convert.ToHexString(
System.Security.Cryptography.MD5.HashData(
System.Text.Encoding.UTF8.GetBytes(dataSb.ToString())
)
);
var sampleDataSet = sampleRecords _logger.LogInformation("Хэш калибровочных записей: {}", dataHashCode);
.AsEnumerable()
.Where(r => r.Features.Count() > 0 && r.Features.Count() % FEATURES_LENGTH == 0) Directory.CreateDirectory(".models_cache");
.SelectMany(r => r.Features.Chunk(FEATURES_LENGTH)) var cachedModelPath = Path.Combine(".models_cache", dataHashCode) + ".zip";
.Select(f => new ClusterizationData
if (File.Exists(cachedModelPath))
{
_logger.LogInformation("Модель найдена к кэше");
var model = _ml.Model.Load(cachedModelPath, out _);
_clusterizationEngine = _ml.Model.CreatePredictionEngine<ClusterizationData, ClusterizationPrediction>(model);
}
else
{
_logger.LogInformation("Найдено калибровочных записей проб для марки: {}", sampleRecords.Count());
if (sampleRecords.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных записей для марки: \"{brand.BrandName} - {brand.MixtureName}\"");
cancellationToken.ThrowIfCancellationRequested();
if (separator.Features.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных признаков для сепаратора: \"{separator.Type}\"");
if (separator.Features.Count() % FEATURES_LENGTH != 0)
throw new InvalidDataException($"Неожидаемое количество калибровочных признаков для сепаратора: \"{separator.Type}\"");
var separatorDataSet = separator.Features.Chunk(FEATURES_LENGTH)
.Select(v => new ClusterizationData
{
Label = "separator",
Features = v.ToArray()
})
.Shuffle();
_logger.LogInformation("Загружено векторов для сепаратора: {}", separatorDataSet.Count());
if (separatorDataSet.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных данных для сепаратора: \"{separator.Type}\"");
cancellationToken.ThrowIfCancellationRequested();
foreach (var sampleRecord in sampleRecords)
{ {
Label = "sample", if (sampleRecord.Features.Count() == 0)
Features = f.ToArray() _logger.LogError("Не найдено калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
}) if (sampleRecord.Features.Count() % FEATURES_LENGTH != 0)
.Shuffle(); _logger.LogError("Неожидаемое количество калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
}
_logger.LogInformation("Загружено векторов для пробы: {}", sampleDataSet.Count()); var sampleDataSet = sampleRecords
if (separatorDataSet.Count() == 0) .AsEnumerable()
throw new InvalidDataException($"Не найдено калибровочных данных для марки: \"{brand.BrandName} - {brand.MixtureName}\""); .Where(r => r.Features.Count() > 0 && r.Features.Count() % FEATURES_LENGTH == 0)
cancellationToken.ThrowIfCancellationRequested(); .SelectMany(r => r.Features.Chunk(FEATURES_LENGTH))
.Select(f => new ClusterizationData
{
Label = "sample",
Features = f.ToArray()
})
.Shuffle();
var trainDataSet = Enumerable.Concat(separatorDataSet, sampleDataSet); _logger.LogInformation("Загружено векторов для пробы: {}", sampleDataSet.Count());
_logger.LogInformation("Загружено векторов для обучения: {}", trainDataSet.Count()); if (separatorDataSet.Count() == 0)
var trainData = _ml.Data.LoadFromEnumerable(trainDataSet); throw new InvalidDataException($"Не найдено калибровочных данных для марки: \"{brand.BrandName} - {brand.MixtureName}\"");
cancellationToken.ThrowIfCancellationRequested(); cancellationToken.ThrowIfCancellationRequested();
_logger.LogInformation("Создание конвейера"); var trainDataSet = Enumerable.Concat(separatorDataSet, sampleDataSet);
var pipeline = _ml.Transforms.Conversion.MapValueToKey("Label") _logger.LogInformation("Загружено векторов для обучения: {}", trainDataSet.Count());
.Append(_ml.MulticlassClassification.Trainers.LbfgsMaximumEntropy()) var trainData = _ml.Data.LoadFromEnumerable(trainDataSet);
.Append(_ml.Transforms.Conversion.MapKeyToValue("PredictedLabel")); cancellationToken.ThrowIfCancellationRequested();
cancellationToken.ThrowIfCancellationRequested();
_logger.LogInformation("Обучение модели"); _logger.LogInformation("Создание конвейера");
var model = pipeline.Fit(trainData); var pipeline = _ml.Transforms.Conversion.MapValueToKey("Label")
cancellationToken.ThrowIfCancellationRequested(); .Append(_ml.MulticlassClassification.Trainers.LbfgsMaximumEntropy())
.Append(_ml.Transforms.Conversion.MapKeyToValue("PredictedLabel"));
cancellationToken.ThrowIfCancellationRequested();
_clusterizationEngine = _ml.Model.CreatePredictionEngine<ClusterizationData, ClusterizationPrediction>(model); _logger.LogInformation("Обучение модели");
model.Dispose(); var model = pipeline.Fit(trainData);
cancellationToken.ThrowIfCancellationRequested();
_ml.Model.Save(model, trainData.Schema, cachedModelPath);
_clusterizationEngine = _ml.Model.CreatePredictionEngine<ClusterizationData, ClusterizationPrediction>(model);
model.Dispose();
}
} }
private async Task InitializeMlRegressionEngine(BrandRecord brand, CancellationToken cancellationToken) private async Task InitializeMlRegressionEngine(BrandRecord brand, CancellationToken cancellationToken)
{ {
@@ -217,79 +250,110 @@ public partial class AnalyzerService
var sampleRecords = _context.SampleRecords var sampleRecords = _context.SampleRecords
.Where(r => r.BrandId == brand.Id); .Where(r => r.BrandId == brand.Id);
_logger.LogInformation("Найдено калибровочных записей проб для марки: {}", sampleRecords.Count()); var dataSb = new System.Text.StringBuilder();
if (sampleRecords.Count() == 0) dataSb.AppendLine(brand.Id.ToString());
throw new InvalidDataException($"Не найдено калибровочных записей для марки: \"{brand.BrandName} - {brand.MixtureName}\""); dataSb.AppendLine(brand.EditDateTime.ToString());
cancellationToken.ThrowIfCancellationRequested(); foreach (var record in sampleRecords.OrderBy(r => r.Id))
foreach (var sampleRecord in sampleRecords)
{ {
if (sampleRecord.Features.Count() == 0) dataSb.AppendLine(record.Id.ToString());
_logger.LogError("Не найдено калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName); dataSb.AppendLine(record.EditDateTime.ToString());
if (sampleRecord.Features.Count() % FEATURES_LENGTH != 0)
_logger.LogError("Неожидаемое количество калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
} }
var dataHashCode = Convert.ToHexString(
_logger.LogInformation("Подготовка данных"); System.Security.Cryptography.MD5.HashData(
var trainDataSet = sampleRecords System.Text.Encoding.UTF8.GetBytes(dataSb.ToString())
.AsEnumerable()
.Where(r => r.Features.Count() > 0 && r.Features.Count() % FEATURES_LENGTH == 0)
.Where(r => r.MeasuredContent >= 0 || r.MixtureActualRate >= 0)
.SelectMany(r =>
r.Features
.Chunk(FEATURES_LENGTH)
.Select(f =>
new RegressionData
{
Value = (float)(r.MeasuredContent >= 0 ? r.MeasuredContent : r.MixtureActualRate),
Features = f.ToArray()
}
)
) )
.Shuffle() );
.ToList();
_logger.LogInformation("Загружено векторов для обучения: {}", trainDataSet.Count()); _logger.LogInformation("Хэш калибровочных записей: {}", dataHashCode);
if (trainDataSet.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных данных для марки: \"{brand.BrandName} - {brand.MixtureName}\"");
cancellationToken.ThrowIfCancellationRequested();
int rowCount = trainDataSet.Count(); Directory.CreateDirectory(".models_cache");
double[][] featuresArray = trainDataSet.Select(d => d.Features.Select(f => (double)f).ToArray()).ToArray(); var cachedModelPath = Path.Combine(".models_cache", dataHashCode) + ".zip";
Matrix<double> matrix = DenseMatrix.OfRowArrays(featuresArray);
for (int i = 0; i < IMAGES_COUNT; i++) if (File.Exists(cachedModelPath))
{ {
matrix = matrix.RemoveColumn(i * 3); _logger.LogInformation("Модель найдена к кэше");
matrix = matrix.RemoveColumn(i * 3); var model = _ml.Model.Load(cachedModelPath, out _);
matrix = matrix.RemoveColumn(i * 3); _regressionEngine = _ml.Model.CreatePredictionEngine<RegressionData, RegressionPrediction>(model);
} }
cancellationToken.ThrowIfCancellationRequested(); else
{
_logger.LogInformation("Найдено калибровочных записей проб для марки: {}", sampleRecords.Count());
if (sampleRecords.Count() == 0)
throw new InvalidDataException($"Не найдено калибровочных записей для марки: \"{brand.BrandName} - {brand.MixtureName}\"");
cancellationToken.ThrowIfCancellationRequested();
Vector<double> meanVector = matrix.ColumnSums() / matrix.RowCount; foreach (var sampleRecord in sampleRecords)
Matrix<double> covarianceMatrix = MatrixCovariance(matrix, meanVector);
Matrix<double> invCovarianceMatrix = covarianceMatrix.Inverse();
trainDataSet = trainDataSet
.Where((data, index) =>
{ {
var row = matrix.Row(index); if (sampleRecord.Features.Count() == 0)
double distance = CalculateMahalanobis(row, meanVector, invCovarianceMatrix); _logger.LogError("Не найдено калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
if (sampleRecord.Features.Count() % FEATURES_LENGTH != 0)
_logger.LogError("Неожидаемое количество калибровочных признаков для пробы: \"{}\"", sampleRecord.SampleName);
}
return distance <= FILTERING_THRESHOLD; _logger.LogInformation("Подготовка данных");
}) var trainDataSet = sampleRecords
.ToList(); .AsEnumerable()
_logger.LogInformation("Векторов для обучения после фильтрации: {}", trainDataSet.Count()); .Where(r => r.Features.Count() > 0 && r.Features.Count() % FEATURES_LENGTH == 0)
var trainData = _ml.Data.LoadFromEnumerable(trainDataSet); .Where(r => r.MeasuredContent >= 0 || r.MixtureActualRate >= 0)
cancellationToken.ThrowIfCancellationRequested(); .SelectMany(r =>
r.Features
.Chunk(FEATURES_LENGTH)
.Select(f =>
new RegressionData
{
Value = (float)(r.MeasuredContent >= 0 ? r.MeasuredContent : r.MixtureActualRate),
Features = f.ToArray()
}
)
)
.Shuffle()
.ToList();
_logger.LogInformation("Создание конвейера"); _logger.LogInformation("Загружено векторов для обучения: {}", trainDataSet.Count());
var pipeline = _ml.Regression.Trainers.FastTree();
cancellationToken.ThrowIfCancellationRequested();
_logger.LogInformation("Обучение модели"); if (trainDataSet.Count() == 0)
var model = pipeline.Fit(trainData); throw new InvalidDataException($"Не найдено калибровочных данных для марки: \"{brand.BrandName} - {brand.MixtureName}\"");
_regressionEngine = _ml.Model.CreatePredictionEngine<RegressionData, RegressionPrediction>(model); cancellationToken.ThrowIfCancellationRequested();
model.Dispose(); int rowCount = trainDataSet.Count();
double[][] featuresArray = trainDataSet.Select(d => d.Features.Select(f => (double)f).ToArray()).ToArray();
Matrix<double> matrix = DenseMatrix.OfRowArrays(featuresArray);
for (int i = 0; i < IMAGES_COUNT; i++)
{
matrix = matrix.RemoveColumn(i * 3);
matrix = matrix.RemoveColumn(i * 3);
matrix = matrix.RemoveColumn(i * 3);
}
cancellationToken.ThrowIfCancellationRequested();
Vector<double> meanVector = matrix.ColumnSums() / matrix.RowCount;
Matrix<double> covarianceMatrix = MatrixCovariance(matrix, meanVector);
Matrix<double> invCovarianceMatrix = covarianceMatrix.Inverse();
trainDataSet = trainDataSet
.Where((data, index) =>
{
var row = matrix.Row(index);
double distance = CalculateMahalanobis(row, meanVector, invCovarianceMatrix);
return distance <= FILTERING_THRESHOLD;
})
.ToList();
_logger.LogInformation("Векторов для обучения после фильтрации: {}", trainDataSet.Count());
var trainData = _ml.Data.LoadFromEnumerable(trainDataSet);
cancellationToken.ThrowIfCancellationRequested();
_logger.LogInformation("Создание конвейера");
var pipeline = _ml.Regression.Trainers.FastTree();
cancellationToken.ThrowIfCancellationRequested();
_logger.LogInformation("Обучение модели");
var model = pipeline.Fit(trainData);
_ml.Model.Save(model, trainData.Schema, cachedModelPath);
_regressionEngine = _ml.Model.CreatePredictionEngine<RegressionData, RegressionPrediction>(model);
model.Dispose();
}
} }
public async Task<(Mat avgImage, Rect roi, Mat mask, IEnumerable<Point[]> contours, Mat result)?> Analyze(IEnumerable<ImageData> imagesData) public async Task<(Mat avgImage, Rect roi, Mat mask, IEnumerable<Point[]> contours, Mat result)?> Analyze(IEnumerable<ImageData> imagesData)