diff --git a/src/Machine/src/Serval.Machine.Shared/Services/BuildJob.cs b/src/Machine/src/Serval.Machine.Shared/Services/BuildJob.cs index f2f877db..60d4aeb3 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/BuildJob.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/BuildJob.cs @@ -13,10 +13,11 @@ public virtual Task RunAsync( string engineId, string buildId, string? buildOptions, + string? model, CancellationToken cancellationToken ) { - return RunAsync(engineId, buildId, null, buildOptions, cancellationToken); + return RunAsync(engineId, buildId, null, buildOptions, model, cancellationToken); } } @@ -40,6 +41,7 @@ public virtual async Task RunAsync( string buildId, TData data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -53,7 +55,7 @@ CancellationToken cancellationToken return; } - await DoWorkAsync(engineId, buildId, data, buildOptions, cancellationToken); + await DoWorkAsync(engineId, buildId, data, buildOptions, model, cancellationToken); } catch (OperationCanceledException e) { @@ -143,6 +145,7 @@ protected abstract Task DoWorkAsync( string buildId, TData data, string? buildOptions, + string? model, CancellationToken cancellationToken ); diff --git a/src/Machine/src/Serval.Machine.Shared/Services/BuildJobService.cs b/src/Machine/src/Serval.Machine.Shared/Services/BuildJobService.cs index 0ee28ff6..fcb40dc4 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/BuildJobService.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/BuildJobService.cs @@ -80,7 +80,6 @@ public async Task StartBuildJobAsync( stage, data, buildOptions, - model, cancellationToken ); try @@ -97,6 +96,7 @@ public async Task StartBuildJobAsync( ) ), u => + { u.Set( e => e.CurrentBuild, new Build @@ -112,7 +112,8 @@ public async Task StartBuildJobAsync( JobData = jobData, ExecutionData = new BuildExecutionData(), } - ), + ); + }, cancellationToken: cancellationToken ); if (engine is null) diff --git a/src/Machine/src/Serval.Machine.Shared/Services/ClearMLBuildJobRunner.cs b/src/Machine/src/Serval.Machine.Shared/Services/ClearMLBuildJobRunner.cs index 825c70fa..a4211047 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/ClearMLBuildJobRunner.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/ClearMLBuildJobRunner.cs @@ -40,7 +40,6 @@ public async Task DeleteEngineAsync(string engineId, CancellationToken cancellat BuildStage stage, object? data = null, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ) { @@ -58,7 +57,6 @@ public async Task DeleteEngineAsync(string engineId, CancellationToken cancellat _options[engineType].ModelType, stage, buildOptions, - model, cancellationToken ); string jobId = await _clearMLService.CreateTaskAsync( diff --git a/src/Machine/src/Serval.Machine.Shared/Services/ClearMLMonitorService.cs b/src/Machine/src/Serval.Machine.Shared/Services/ClearMLMonitorService.cs index 3a825302..60469fc6 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/ClearMLMonitorService.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/ClearMLMonitorService.cs @@ -194,6 +194,7 @@ await UpdateTrainJobStatus( (int)GetMetric(task, SummaryMetric, TrainCorpusSizeVariant), GetMetric(task, SummaryMetric, ConfidenceVariant), engine.CurrentBuild.Options, + engine.CurrentBuild.Model, cancellationToken ); if (canceling) @@ -285,6 +286,7 @@ private async Task TrainJobCompletedAsync( int corpusSize, double confidence, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -298,6 +300,7 @@ CancellationToken cancellationToken BuildStage.Postprocess, (corpusSize, confidence), buildOptions, + model, cancellationToken: cancellationToken ); } diff --git a/src/Machine/src/Serval.Machine.Shared/Services/IBuildJobRunner.cs b/src/Machine/src/Serval.Machine.Shared/Services/IBuildJobRunner.cs index afaaf5e3..c9b5b37d 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/IBuildJobRunner.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/IBuildJobRunner.cs @@ -15,7 +15,6 @@ public interface IBuildJobRunner BuildStage stage, object? data = null, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ); diff --git a/src/Machine/src/Serval.Machine.Shared/Services/IClearMLBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Shared/Services/IClearMLBuildJobFactory.cs index 0c4c34a6..3dbd6e2b 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/IClearMLBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/IClearMLBuildJobFactory.cs @@ -10,7 +10,6 @@ Task CreateJobScriptAsync( string modelType, BuildStage stage, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ); } diff --git a/src/Machine/src/Serval.Machine.Shared/Services/ILocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Shared/Services/ILocalBuildJobFactory.cs index 55cc1d4a..d5aa506c 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/ILocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/ILocalBuildJobFactory.cs @@ -13,6 +13,7 @@ Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ); } diff --git a/src/Machine/src/Serval.Machine.Shared/Services/LocalBuildJobRunner.cs b/src/Machine/src/Serval.Machine.Shared/Services/LocalBuildJobRunner.cs index 8be14691..d3c55e96 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/LocalBuildJobRunner.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/LocalBuildJobRunner.cs @@ -44,7 +44,6 @@ public Task CreateEngineAsync( BuildStage stage, object? data = null, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ) { @@ -215,7 +214,7 @@ private void EnqueueRecoveredJob(string engineId, Build build, EngineType engine if ( _pendingJobs.TryAdd( build.JobId, - new JobInfo(engineId, build.BuildId, engineType, build.Stage, build.JobData, build.Options) + new JobInfo(engineId, build.BuildId, engineType, build.Stage, build.JobData, build.Options, build.Model) ) ) { @@ -250,7 +249,8 @@ CancellationToken cancellationToken engine.Type, build.Stage, build.JobData, - build.Options + build.Options, + build.Model ) ) ) @@ -316,6 +316,7 @@ await factory.RunAsync( info.Stage, info.JobData, info.BuildOptions, + info.Model, cts.Token ); } @@ -336,6 +337,7 @@ private record JobInfo( EngineType EngineType, BuildStage Stage, string? JobData, - string? BuildOptions + string? BuildOptions, + string? Model ); } diff --git a/src/Machine/src/Serval.Machine.Shared/Services/PreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Shared/Services/PreprocessBuildJob.cs index ac4d5ebb..dfdf2f6d 100644 --- a/src/Machine/src/Serval.Machine.Shared/Services/PreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Shared/Services/PreprocessBuildJob.cs @@ -40,6 +40,7 @@ protected override async Task DoWorkAsync( string buildId, IReadOnlyList data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -57,6 +58,7 @@ await UpdateBuildExecutionData( engine.SourceLanguage, engine.TargetLanguage, isNonPersistedTranslationEngine, + model ?? "unknown", data, cancellationToken ); @@ -79,6 +81,7 @@ await UpdateBuildExecutionData( buildId, BuildStage.Train, buildOptions: buildOptions, + model: model, cancellationToken: cancellationToken ); if (canceling) @@ -92,6 +95,7 @@ protected abstract Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ); @@ -134,8 +138,8 @@ protected virtual IReadOnlyList GetDiagnostics( bool sourceLanguageHasNativeSupport, bool targetLanguageHasNativeSupport, bool isNonPersistedTranslationEngine, - string modelName, - IReadOnlyList parallelCorpora + IReadOnlyList parallelCorpora, + string? modelName = null ) { List diagnostics = []; diff --git a/src/Machine/src/Serval.Machine.Translation/Models/BaseModels.cs b/src/Machine/src/Serval.Machine.Translation/Models/BaseModels.cs deleted file mode 100644 index 22b1fb67..00000000 --- a/src/Machine/src/Serval.Machine.Translation/Models/BaseModels.cs +++ /dev/null @@ -1,8 +0,0 @@ -namespace Serval.Machine.Translation.Models; - -public static class Models -{ - public const string Nllb = "NLLB"; - public const string Nllb600m = "NLLB600m"; - public const string NllbTesting = "NLLBTesting"; -} diff --git a/src/Machine/src/Serval.Machine.Translation/Models/Models.cs b/src/Machine/src/Serval.Machine.Translation/Models/Models.cs new file mode 100644 index 00000000..798014ff --- /dev/null +++ b/src/Machine/src/Serval.Machine.Translation/Models/Models.cs @@ -0,0 +1,8 @@ +namespace Serval.Machine.Translation.Models; + +public static class Models +{ + public const string Nllb = "nllb"; + public const string Nllb600m = "nllb-600m"; + public const string NllbTesting = "nllb-testing"; +} diff --git a/src/Machine/src/Serval.Machine.Translation/Services/EchoLocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Translation/Services/EchoLocalBuildJobFactory.cs index 5152b555..f0868e12 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/EchoLocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/EchoLocalBuildJobFactory.cs @@ -30,6 +30,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -38,11 +39,11 @@ CancellationToken cancellationToken case BuildStage.Preprocess: var preprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); var corpora = JsonSerializer.Deserialize>(jobData!, SerializerOptions)!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, model, cancellationToken); break; case BuildStage.Train: var trainJob = ActivatorUtilities.CreateInstance(serviceProvider); - await trainJob.RunAsync(engineId, buildId, null, buildOptions, cancellationToken); + await trainJob.RunAsync(engineId, buildId, null, buildOptions, model, cancellationToken); break; case BuildStage.Postprocess: var postprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); @@ -52,6 +53,7 @@ await postprocessJob.RunAsync( buildId, (postData.TrainCount, postData.Confidence), buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/src/Serval.Machine.Translation/Services/EchoPostprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/EchoPostprocessBuildJob.cs index bbe9e639..665b1065 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/EchoPostprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/EchoPostprocessBuildJob.cs @@ -24,6 +24,7 @@ protected override async Task DoWorkAsync( string buildId, (int, double) data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.Translation/Services/EchoPreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/EchoPreprocessBuildJob.cs index d64995b6..454f629b 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/EchoPreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/EchoPreprocessBuildJob.cs @@ -33,13 +33,11 @@ protected override async Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ) { - string modelName = - (await Engines.GetAsync(e => e.EngineId == engineId, cancellationToken))?.CurrentBuild?.Model?.ToString() - ?? "Unknown"; IReadOnlyList diagnostics = GetDiagnostics( stats.TrainCount, stats.InferenceCount, @@ -48,7 +46,6 @@ CancellationToken cancellationToken sourceLanguageHasNativeSupport: true, targetLanguageHasNativeSupport: true, isNonPersistedTranslationEngine, - modelName, parallelCorpora ); diff --git a/src/Machine/src/Serval.Machine.Translation/Services/EchoTrainingBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/EchoTrainingBuildJob.cs index 16d39810..05a06c95 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/EchoTrainingBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/EchoTrainingBuildJob.cs @@ -13,6 +13,7 @@ protected override async Task DoWorkAsync( string buildId, object? data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.Translation/Services/EchoTranslationEngineService.cs b/src/Machine/src/Serval.Machine.Translation/Services/EchoTranslationEngineService.cs index 2e334ee8..36951128 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/EchoTranslationEngineService.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/EchoTranslationEngineService.cs @@ -67,7 +67,7 @@ await engines.UpdateAsync( ); } - public async Task StartBuildAsync( + public async Task StartBuildAsync( string engineId, string buildId, IReadOnlyList corpora, @@ -84,12 +84,13 @@ public async Task StartBuildAsync( BuildStage.Preprocess, corpora, options, - model, cancellationToken: cancellationToken ); // If there is a pending/running build, then no need to start a new one. if (building) throw new ConflictException(); + + return new() { Model = model }; } public Task CancelBuildAsync(string engineId, CancellationToken cancellationToken = default) => diff --git a/src/Machine/src/Serval.Machine.Translation/Services/NmtClearMLBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Translation/Services/NmtClearMLBuildJobFactory.cs index 4b0d29fe..0aece6c6 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/NmtClearMLBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/NmtClearMLBuildJobFactory.cs @@ -18,7 +18,6 @@ public async Task CreateJobScriptAsync( string modelType, BuildStage stage, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ) { @@ -33,19 +32,6 @@ public async Task CreateJobScriptAsync( string folder = sharedFileUri.GetComponents(UriComponents.Path, UriFormat.Unescaped); _languageTagService.ConvertToFlores200Code(engine.SourceLanguage, out string srcLang); _languageTagService.ConvertToFlores200Code(engine.TargetLanguage, out string trgLang); - if (buildOptions != null && model != null) - { - try - { - JsonNode? buildOptionsJsonNode = JsonNode.Parse(buildOptions); - if (buildOptionsJsonNode != null && buildOptionsJsonNode is JsonObject buildOptionsJsonObject) - buildOptionsJsonObject["parent_model_name"] = GetFullModelName(model); - } - catch (Exception e) - { - throw new InvalidOperationException($"Unable to parse field build options : {e.Message}", e); - } - } return "from machine.jobs.build_nmt_engine import run\n" + "args = {\n" + $" 'model_type': '{modelType}',\n" @@ -68,15 +54,4 @@ public async Task CreateJobScriptAsync( throw new ArgumentException("Unknown build stage.", nameof(stage)); } } - - private static string GetFullModelName(string model) - { - return model switch - { - Models.Models.Nllb => "facebook/nllb-200-distilled-1.3B", - Models.Models.Nllb600m => "facebook/nllb-200-distilled-600M", - Models.Models.NllbTesting => "hf-internal-testing/tiny-random-nllb", - _ => throw new ArgumentException($"Unknown base model {model}."), - }; - } } diff --git a/src/Machine/src/Serval.Machine.Translation/Services/NmtEngineService.cs b/src/Machine/src/Serval.Machine.Translation/Services/NmtEngineService.cs index 586d8fce..caf2843d 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/NmtEngineService.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/NmtEngineService.cs @@ -17,6 +17,20 @@ ISharedFileService sharedFileService private readonly ISharedFileService _sharedFileService = sharedFileService; public const string ModelDirectory = "models/"; + private static readonly IReadOnlyDictionary ModelToFullModelName = new Dictionary() + { + [Models.Models.Nllb] = "facebook/nllb-200-distilled-1.3B", + [Models.Models.Nllb600m] = "facebook/nllb-200-distilled-600M", + [Models.Models.NllbTesting] = "hf-internal-testing/tiny-random-nllb", + }; + + private static readonly IReadOnlyDictionary FullModelNameToModel = new Dictionary() + { + ["facebook/nllb-200-distilled-1.3B"] = Models.Models.Nllb, + ["facebook/nllb-200-distilled-600M"] = Models.Models.Nllb600m, + ["hf-internal-testing/tiny-random-nllb"] = Models.Models.NllbTesting, + }; + public static string GetModelPath(string engineId, int buildRevision) { return $"{ModelDirectory}{engineId}_{buildRevision}.tar.gz"; @@ -80,7 +94,7 @@ await _engines.UpdateAsync( ); } - public async Task StartBuildAsync( + public async Task StartBuildAsync( string engineId, string buildId, IReadOnlyList corpora, @@ -89,6 +103,36 @@ public async Task StartBuildAsync( CancellationToken cancellationToken = default ) { + JsonObject? buildOptionsJsonObject = []; + if (options != null) + { + try + { + JsonNode? buildOptionsJsonNode = JsonNode.Parse(options); + if (buildOptionsJsonNode is JsonObject obj) + buildOptionsJsonObject = obj; + } + catch (Exception e) + { + throw new InvalidOperationException($"Unable to parse field build options : {e.Message}", e); + } + } + + if ( + buildOptionsJsonObject.ContainsKey("parent_model_name") + && buildOptionsJsonObject["parent_model_name"] != null + && model == null + ) + { + model = GetModelName(buildOptionsJsonObject["parent_model_name"]!.GetValue()); + } + else + { + model ??= Models.Models.Nllb; + buildOptionsJsonObject["parent_model_name"] = GetFullModelName(model); + options = buildOptionsJsonObject.ToJsonString(); + } + bool building = !await _buildJobService.StartBuildJobAsync( BuildJobRunnerType.Local, EngineType.Nmt, @@ -103,6 +147,8 @@ public async Task StartBuildAsync( // If there is a pending/running build, then no need to start a new one. if (building) throw new ConflictException(); + + return new() { Model = model }; } public Task CancelBuildAsync(string engineId, CancellationToken cancellationToken = default) @@ -208,4 +254,18 @@ private async Task GetEngineAsync(string engineId, Cancellati throw new EngineNotFoundException($"The engine {engineId} does not exist."); return engine; } + + private static string GetFullModelName(string model) + { + if (ModelToFullModelName.TryGetValue(model, out string? fullModelName) && fullModelName != null) + return fullModelName; + throw new InvalidOperationException($"Unknown model {model}."); + } + + private static string GetModelName(string fullModelName) + { + if (FullModelNameToModel.TryGetValue(fullModelName, out string? model) && model != null) + return model; + throw new InvalidOperationException($"Unknown full model name {fullModelName}."); + } } diff --git a/src/Machine/src/Serval.Machine.Translation/Services/NmtLocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Translation/Services/NmtLocalBuildJobFactory.cs index 60f0679c..2ece5826 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/NmtLocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/NmtLocalBuildJobFactory.cs @@ -29,6 +29,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -37,7 +38,7 @@ CancellationToken cancellationToken case BuildStage.Preprocess: var preprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); var corpora = JsonSerializer.Deserialize>(jobData!, SerializerOptions)!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, model, cancellationToken); break; case BuildStage.Postprocess: var postprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); @@ -47,6 +48,7 @@ await postprocessJob.RunAsync( buildId, (postData.TrainCount, postData.Confidence), buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/src/Serval.Machine.Translation/Services/NmtPreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/NmtPreprocessBuildJob.cs index b2ce81aa..c6016f24 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/NmtPreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/NmtPreprocessBuildJob.cs @@ -58,6 +58,7 @@ protected override async Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ) @@ -65,9 +66,6 @@ CancellationToken cancellationToken bool sourceLanguageHasNativeSupport = ResolveLanguageCode(sourceLanguageTag, out string resolvedSourceLanguage); bool targetLanguageHasNativeSupport = ResolveLanguageCode(targetLanguageTag, out string resolvedTargetLanguage); - string modelName = - (await Engines.GetAsync(e => e.EngineId == engineId, cancellationToken))?.CurrentBuild?.Model?.ToString() - ?? "Unknown"; IReadOnlyList diagnostics = GetDiagnostics( stats.TrainCount, stats.InferenceCount, @@ -76,8 +74,8 @@ CancellationToken cancellationToken sourceLanguageHasNativeSupport, targetLanguageHasNativeSupport, isNonPersistedTranslationEngine, - modelName, - parallelCorpora + parallelCorpora, + modelName ); IReadOnlyList warnings = diagnostics.Select(d => d.Message).ToList(); @@ -147,10 +145,12 @@ protected override IReadOnlyList GetDiagnostics( bool sourceLanguageHasNativeSupport, bool targetLanguageHasNativeSupport, bool isNonPersistedTranslationEngine, - string modelName, - IReadOnlyList parallelCorpora + IReadOnlyList parallelCorpora, + string? modelName = null ) { + modelName ??= "unknown"; + List diagnostics = []; // Has at least a Gospel of Mark amount of data and not the special case of no data which will be caught elsewhere @@ -222,8 +222,8 @@ .. base.GetDiagnostics( sourceLanguageHasNativeSupport, targetLanguageHasNativeSupport, isNonPersistedTranslationEngine, - modelName, - parallelCorpora + parallelCorpora, + modelName ), .. diagnostics, ]; diff --git a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferClearMLBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferClearMLBuildJobFactory.cs index 85c3ec26..5cfe488f 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferClearMLBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferClearMLBuildJobFactory.cs @@ -16,7 +16,6 @@ public async Task CreateJobScriptAsync( string modelType, BuildStage stage, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ) { diff --git a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferEngineService.cs b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferEngineService.cs index 20793653..609e8d92 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferEngineService.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferEngineService.cs @@ -179,7 +179,7 @@ await _trainSegmentPairs.InsertAsync( state.Touch(); } - public async Task StartBuildAsync( + public async Task StartBuildAsync( string engineId, string buildId, IReadOnlyList corpora, @@ -196,7 +196,6 @@ public async Task StartBuildAsync( BuildStage.Preprocess, corpora, options, - model, cancellationToken: cancellationToken ); // If there is a pending/running build, then no need to start a new one. @@ -205,6 +204,8 @@ public async Task StartBuildAsync( SmtTransferEngineState state = _stateService.Get(engineId); state.Touch(); + + return new(); } public async Task CancelBuildAsync(string engineId, CancellationToken cancellationToken = default) diff --git a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferLocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferLocalBuildJobFactory.cs index 835a9817..a8a1744a 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferLocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/SmtTransferLocalBuildJobFactory.cs @@ -30,6 +30,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -38,7 +39,7 @@ CancellationToken cancellationToken case BuildStage.Preprocess: var preprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); var corpora = JsonSerializer.Deserialize>(jobData!, SerializerOptions)!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, model, cancellationToken); break; case BuildStage.Postprocess: var postprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); @@ -48,6 +49,7 @@ await postprocessJob.RunAsync( buildId, (postData.TrainCount, postData.Confidence), buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/src/Serval.Machine.Translation/Services/TranslationPostprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/TranslationPostprocessBuildJob.cs index cfe2f76e..158291e6 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/TranslationPostprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/TranslationPostprocessBuildJob.cs @@ -24,6 +24,7 @@ protected override async Task DoWorkAsync( string buildId, (int, double) data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.Translation/Services/TranslationPreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.Translation/Services/TranslationPreprocessBuildJob.cs index 790cd1d2..a809f77f 100644 --- a/src/Machine/src/Serval.Machine.Translation/Services/TranslationPreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.Translation/Services/TranslationPreprocessBuildJob.cs @@ -92,13 +92,11 @@ protected override async Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ) { - string modelName = - (await Engines.GetAsync(e => e.EngineId == engineId, cancellationToken))?.CurrentBuild?.Model?.ToString() - ?? "Unknown"; IReadOnlyList diagnostics = GetDiagnostics( stats.TrainCount, stats.InferenceCount, @@ -107,8 +105,8 @@ CancellationToken cancellationToken true, true, isNonPersistedTranslationEngine, - modelName, - parallelCorpora + parallelCorpora, + modelName ); IReadOnlyList warnings = diagnostics.Select(d => d.Message).ToList(); diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentLocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentLocalBuildJobFactory.cs index 5854fcb7..74bc9567 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentLocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentLocalBuildJobFactory.cs @@ -30,6 +30,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -40,11 +41,11 @@ CancellationToken cancellationToken serviceProvider ); var corpora = JsonSerializer.Deserialize>(jobData!, SerializerOptions)!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, model, cancellationToken); break; case BuildStage.Train: var trainJob = ActivatorUtilities.CreateInstance(serviceProvider); - await trainJob.RunAsync(engineId, buildId, null, buildOptions, cancellationToken); + await trainJob.RunAsync(engineId, buildId, null, buildOptions, model, cancellationToken); break; case BuildStage.Postprocess: var postprocessJob = ActivatorUtilities.CreateInstance( @@ -56,6 +57,7 @@ await postprocessJob.RunAsync( buildId, (postData.TrainCount, postData.Confidence), buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPostprocessBuildJob.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPostprocessBuildJob.cs index 29edc307..2e15362f 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPostprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPostprocessBuildJob.cs @@ -24,6 +24,7 @@ protected override async Task DoWorkAsync( string buildId, (int, double) data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPreprocessBuildJob.cs index bd847597..b9341304 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentPreprocessBuildJob.cs @@ -93,13 +93,11 @@ protected override async Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ) { - string modelName = - (await Engines.GetAsync(e => e.EngineId == engineId, cancellationToken))?.CurrentBuild?.Model?.ToString() - ?? "Unknown"; IReadOnlyList diagnostics = GetDiagnostics( stats.TrainCount, stats.InferenceCount, @@ -108,7 +106,6 @@ CancellationToken cancellationToken sourceLanguageHasNativeSupport: true, targetLanguageHasNativeSupport: true, isNonPersistedTranslationEngine, - modelName, parallelCorpora ); diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentTrainingBuildJob.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentTrainingBuildJob.cs index aa8643f3..2a994a94 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentTrainingBuildJob.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/EchoWordAlignmentTrainingBuildJob.cs @@ -13,6 +13,7 @@ protected override async Task DoWorkAsync( string buildId, object? data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalClearMLBuildJobFactory.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalClearMLBuildJobFactory.cs index e593679e..60a7303c 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalClearMLBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalClearMLBuildJobFactory.cs @@ -16,7 +16,6 @@ public async Task CreateJobScriptAsync( string modelType, BuildStage stage, string? buildOptions = null, - string? model = null, CancellationToken cancellationToken = default ) { diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalLocalBuildJobFactory.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalLocalBuildJobFactory.cs index e5d736be..62bb8d96 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalLocalBuildJobFactory.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalLocalBuildJobFactory.cs @@ -30,6 +30,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -38,7 +39,7 @@ CancellationToken cancellationToken case BuildStage.Preprocess: var preprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); var corpora = JsonSerializer.Deserialize>(jobData!, SerializerOptions)!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, model, cancellationToken); break; case BuildStage.Postprocess: var postprocessJob = ActivatorUtilities.CreateInstance(serviceProvider); @@ -48,6 +49,7 @@ await postprocessJob.RunAsync( buildId, (postData.TrainCount, postData.Confidence), buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalPostprocessBuildJob.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalPostprocessBuildJob.cs index 77fe1d89..68190c64 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalPostprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/StatisticalPostprocessBuildJob.cs @@ -31,6 +31,7 @@ protected override async Task DoWorkAsync( string buildId, (int, double) data, string? buildOptions, + string? model, CancellationToken cancellationToken ) { diff --git a/src/Machine/src/Serval.Machine.WordAlignment/Services/WordAlignmentPreprocessBuildJob.cs b/src/Machine/src/Serval.Machine.WordAlignment/Services/WordAlignmentPreprocessBuildJob.cs index 3b61b9bf..b4a77a18 100644 --- a/src/Machine/src/Serval.Machine.WordAlignment/Services/WordAlignmentPreprocessBuildJob.cs +++ b/src/Machine/src/Serval.Machine.WordAlignment/Services/WordAlignmentPreprocessBuildJob.cs @@ -91,13 +91,11 @@ protected override async Task UpdateBuildExecutionData( string sourceLanguageTag, string targetLanguageTag, bool isNonPersistedTranslationEngine, + string modelName, IReadOnlyList parallelCorpora, CancellationToken cancellationToken ) { - string modelName = - (await Engines.GetAsync(e => e.EngineId == engineId, cancellationToken))?.CurrentBuild?.Model?.ToString() - ?? "Unknown"; IReadOnlyList diagnostics = GetDiagnostics( stats.TrainCount, stats.InferenceCount, @@ -106,8 +104,7 @@ CancellationToken cancellationToken sourceLanguageHasNativeSupport: true, targetLanguageHasNativeSupport: true, isNonPersistedTranslationEngine, - modelName, - parallelCorpora + parallelCorpora: parallelCorpora ); IReadOnlyList warnings = diagnostics.Select(d => d.Message).ToList(); @@ -152,7 +149,6 @@ CancellationToken cancellationToken Warnings = warnings, Diagnostics = diagnostics, DiagnosticsTruncated = diagnosticsTruncated, - EngineSourceLanguageTag = sourceLanguageTag, EngineTargetLanguageTag = targetLanguageTag, }; diff --git a/src/Machine/test/Serval.Machine.Translation.Tests/Services/NmtEngineServiceTests.cs b/src/Machine/test/Serval.Machine.Translation.Tests/Services/NmtEngineServiceTests.cs index aa81249e..66d0a806 100644 --- a/src/Machine/test/Serval.Machine.Translation.Tests/Services/NmtEngineServiceTests.cs +++ b/src/Machine/test/Serval.Machine.Translation.Tests/Services/NmtEngineServiceTests.cs @@ -21,6 +21,128 @@ public async Task StartBuildAsync() }); } + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_DefaultModel() + { + using var env = new TestEnvironment(); + env.PersistModel(); + env.UseInfiniteTrainJob(); + + TranslationEngine engine = env.Engines.Get("engine1"); + await env.Service.StartBuildAsync("engine1", "build1", Array.Empty(), "{}"); + await env.WaitForBuildToStartAsync(); + engine = env.Engines.Get("engine1"); + Assert.That(engine.CurrentBuild, Is.Not.Null); + Assert.That(engine.CurrentBuild.Model, Is.EqualTo(Models.Models.Nllb)); + await env.Service.CancelBuildAsync("engine1"); + await env.WaitForBuildToFinishAsync(); + } + + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_SpecifyModelParameter() + { + using var env = new TestEnvironment(); + env.PersistModel(); + env.UseInfiniteTrainJob(); + + await env.Service.StartBuildAsync( + "engine1", + "build1", + Array.Empty(), + "{}", + Models.Models.NllbTesting + ); + await env.WaitForBuildToStartAsync(); + TranslationEngine engine = env.Engines.Get("engine1"); + Assert.That(engine.CurrentBuild, Is.Not.Null); + Assert.That(engine.CurrentBuild.Model, Is.EqualTo(Models.Models.NllbTesting)); + Assert.That(engine.CurrentBuild.Options, Does.Contain("hf-internal-testing/tiny-random-nllb")); + await env.Service.CancelBuildAsync("engine1"); + await env.WaitForBuildToFinishAsync(); + } + + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_SpecifyModelInOptions() + { + using var env = new TestEnvironment(); + env.PersistModel(); + env.UseInfiniteTrainJob(); + + // Non-default model and also specified in options + await env.Service.StartBuildAsync( + "engine1", + "build1", + Array.Empty(), + "{\"parent_model_name\": \"facebook/nllb-200-distilled-1.3B\"}" + ); + await env.WaitForBuildToStartAsync(); + TranslationEngine engine = env.Engines.Get("engine1"); + Assert.That(engine.CurrentBuild, Is.Not.Null); + Assert.That(engine.CurrentBuild.Model, Is.EqualTo(Models.Models.Nllb)); + Assert.That(engine.CurrentBuild.Options, Does.Contain("facebook/nllb-200-distilled-1.3B")); + await env.Service.CancelBuildAsync("engine1"); + await env.WaitForBuildToFinishAsync(); + } + + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_SpecifyModelInOptionsAndModelParameter() + { + using var env = new TestEnvironment(); + env.PersistModel(); + env.UseInfiniteTrainJob(); + + // Non-default model and also specified in options + await env.Service.StartBuildAsync( + "engine1", + "build1", + Array.Empty(), + "{\"parent_model_name\": \"facebook/nllb-200-distilled-1.3B\"}", + Models.Models.NllbTesting + ); + await env.WaitForBuildToStartAsync(); + TranslationEngine engine = env.Engines.Get("engine1"); + Assert.That(engine.CurrentBuild, Is.Not.Null); + Assert.That(engine.CurrentBuild.Model, Is.EqualTo(Models.Models.NllbTesting)); + Assert.That(engine.CurrentBuild.Options, Does.Contain("hf-internal-testing/tiny-random-nllb")); + await env.Service.CancelBuildAsync("engine1"); + await env.WaitForBuildToFinishAsync(); + } + + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_InvalidModelParameter() + { + using var env = new TestEnvironment(); + env.PersistModel(); + + // Invalid model + Assert.ThrowsAsync(async () => + await env.Service.StartBuildAsync( + "engine1", + "build1", + Array.Empty(), + "{\"parent_model_name\": \"facebook/nllb-200-distilled-1.3B\"}", + "NLLBMisspelled" + ) + ); + } + + [Test] + public async Task StartBuildAsync_ModelIsCorrectlySet_InvalidModelInOptions() + { + using var env = new TestEnvironment(); + env.PersistModel(); + + // Invalid model + Assert.ThrowsAsync(async () => + await env.Service.StartBuildAsync( + "engine1", + "build1", + Array.Empty(), + "{\"parent_model_name\": \"invalid-model\"}" + ) + ); + } + [Test] public async Task CancelBuildAsync_Building() { diff --git a/src/Machine/test/Serval.Machine.Translation.Tests/Services/PreprocessBuildJobTests.cs b/src/Machine/test/Serval.Machine.Translation.Tests/Services/PreprocessBuildJobTests.cs index 0f15b88e..dcfdcb1a 100644 --- a/src/Machine/test/Serval.Machine.Translation.Tests/Services/PreprocessBuildJobTests.cs +++ b/src/Machine/test/Serval.Machine.Translation.Tests/Services/PreprocessBuildJobTests.cs @@ -293,7 +293,7 @@ public async Task RunAsync_BuildWarnings() using (Assert.EnterMultipleScope()) { Assert.That(data["resolvedCode"], Is.Null); - Assert.That(data["modelName"] is string modelName && modelName == "NLLB"); + Assert.That(data["modelName"] is string modelName && modelName == "nllb"); } Assert.That(env.ExecutionData.Diagnostics[10].Code, Is.EqualTo("MODEL-0002")); @@ -302,7 +302,7 @@ public async Task RunAsync_BuildWarnings() using (Assert.EnterMultipleScope()) { Assert.That(data["resolvedCode"], Is.Null); - Assert.That(data["modelName"] is string modelName && modelName == "NLLB"); + Assert.That(data["modelName"] is string modelName && modelName == "nllb"); } env.BuildJobOptions.CurrentValue.Returns(new BuildJobOptions() { MaxWarnings = 1, MaxDiagnostics = 1 }); @@ -348,7 +348,7 @@ public async Task RunAsync_BuildWarnings() using (Assert.EnterMultipleScope()) { Assert.That(data["resolvedCode"], Is.Null); - Assert.That(data["modelName"] is string modelName && modelName == "NLLB"); + Assert.That(data["modelName"] is string modelName && modelName == "nllb"); } Assert.That(env.ExecutionData.Diagnostics[2].Code, Is.EqualTo("MODEL-0002")); @@ -357,7 +357,7 @@ public async Task RunAsync_BuildWarnings() using (Assert.EnterMultipleScope()) { Assert.That(data["resolvedCode"], Is.Null); - Assert.That(data["modelName"] is string modelName && modelName == "NLLB"); + Assert.That(data["modelName"] is string modelName && modelName == "nllb"); } Assert.That(env.ExecutionData.Diagnostics[3].Code, Is.EqualTo("MODEL-0004")); @@ -365,7 +365,7 @@ public async Task RunAsync_BuildWarnings() Assert.That(data, Has.Count.EqualTo(2)); using (Assert.EnterMultipleScope()) { - Assert.That(data["modelName"] is string modelName && modelName == "NLLB"); + Assert.That(data["modelName"] is string modelName && modelName == "nllb"); Assert.That( data["unknownLanguageCodes"] is List unknownLanguageCodes && unknownLanguageCodes.SequenceEqual(["es", "en"]) @@ -934,6 +934,7 @@ public Task RunBuildJobAsync( "build1", corpora.ToList(), useKeyTerms ? null : "{\"use_key_terms\":false}", + "nllb", default ); } diff --git a/src/Machine/test/Serval.Machine.Translation.Tests/Services/SmtTransferEngineServiceTests.cs b/src/Machine/test/Serval.Machine.Translation.Tests/Services/SmtTransferEngineServiceTests.cs index ee59dd71..32d3fe65 100644 --- a/src/Machine/test/Serval.Machine.Translation.Tests/Services/SmtTransferEngineServiceTests.cs +++ b/src/Machine/test/Serval.Machine.Translation.Tests/Services/SmtTransferEngineServiceTests.cs @@ -675,6 +675,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -688,7 +689,14 @@ CancellationToken cancellationToken jobData!, SerializerOptions )!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync( + engineId, + buildId, + corpora, + buildOptions, + model, + cancellationToken + ); break; default: await new SmtTransferLocalBuildJobFactory().RunAsync( @@ -698,6 +706,7 @@ CancellationToken cancellationToken stage, jobData, buildOptions, + model, cancellationToken ); break; diff --git a/src/Machine/test/Serval.Machine.WordAlignment.Tests/Services/StatisticalEngineServiceTests.cs b/src/Machine/test/Serval.Machine.WordAlignment.Tests/Services/StatisticalEngineServiceTests.cs index 303d227c..f556db1f 100644 --- a/src/Machine/test/Serval.Machine.WordAlignment.Tests/Services/StatisticalEngineServiceTests.cs +++ b/src/Machine/test/Serval.Machine.WordAlignment.Tests/Services/StatisticalEngineServiceTests.cs @@ -432,6 +432,7 @@ public async Task RunAsync( BuildStage stage, string? jobData, string? buildOptions, + string? model, CancellationToken cancellationToken ) { @@ -445,7 +446,14 @@ CancellationToken cancellationToken jobData!, SerializerOptions )!; - await preprocessJob.RunAsync(engineId, buildId, corpora, buildOptions, cancellationToken); + await preprocessJob.RunAsync( + engineId, + buildId, + corpora, + buildOptions, + model, + cancellationToken + ); break; default: await new StatisticalLocalBuildJobFactory().RunAsync( @@ -455,6 +463,7 @@ CancellationToken cancellationToken stage, jobData, buildOptions, + model, cancellationToken ); break; diff --git a/src/Serval/src/Serval.Client/Client.g.cs b/src/Serval/src/Serval.Client/Client.g.cs index 92dbf5ec..2024033a 100644 --- a/src/Serval/src/Serval.Client/Client.g.cs +++ b/src/Serval/src/Serval.Client/Client.g.cs @@ -11509,6 +11509,9 @@ public partial class TranslationBuild [Newtonsoft.Json.JsonProperty("options", Required = Newtonsoft.Json.Required.Default, NullValueHandling = Newtonsoft.Json.NullValueHandling.Ignore)] public object? Options { get; set; } = default!; + [Newtonsoft.Json.JsonProperty("model", Required = Newtonsoft.Json.Required.Default, NullValueHandling = Newtonsoft.Json.NullValueHandling.Ignore)] + public string? Model { get; set; } = default!; + [Newtonsoft.Json.JsonProperty("deploymentVersion", Required = Newtonsoft.Json.Required.Default, NullValueHandling = Newtonsoft.Json.NullValueHandling.Ignore)] public string? DeploymentVersion { get; set; } = default!; diff --git a/src/Serval/src/Serval.Shared/Services/BuildDiagnosticService.cs b/src/Serval/src/Serval.Shared/Services/BuildDiagnosticService.cs index 6a2ae62c..236f107d 100644 --- a/src/Serval/src/Serval.Shared/Services/BuildDiagnosticService.cs +++ b/src/Serval/src/Serval.Shared/Services/BuildDiagnosticService.cs @@ -58,6 +58,10 @@ public string FormatMessage(Dictionary data) ["bookId"] = typeof(string), ["modelName"] = typeof(string), }, + DataFormatters = new Dictionary> + { + ["averagePretranslationConfidence"] = obj => ((double)obj).ToString("F3", CultureInfo.InvariantCulture), + }, }, ["MODEL-0004"] = new DiagnosticInfo { diff --git a/src/Serval/src/Serval.Translation.Contracts/ITranslationEngineService.cs b/src/Serval/src/Serval.Translation.Contracts/ITranslationEngineService.cs index eadcf8a4..24ffcbe1 100644 --- a/src/Serval/src/Serval.Translation.Contracts/ITranslationEngineService.cs +++ b/src/Serval/src/Serval.Translation.Contracts/ITranslationEngineService.cs @@ -41,7 +41,7 @@ Task GetModelDownloadUrlAsync( CancellationToken cancellationToken = default ); - Task StartBuildAsync( + Task StartBuildAsync( string engineId, string buildId, IReadOnlyList corpora, diff --git a/src/Serval/src/Serval.Translation.Contracts/StartBuildContract.cs b/src/Serval/src/Serval.Translation.Contracts/StartBuildContract.cs new file mode 100644 index 00000000..f8cd48e1 --- /dev/null +++ b/src/Serval/src/Serval.Translation.Contracts/StartBuildContract.cs @@ -0,0 +1,6 @@ +namespace Serval.Translation.Contracts; + +public record StartBuildContract +{ + public string? Model { get; init; } +} diff --git a/src/Serval/src/Serval.Translation/Dtos/TranslationBuildDto.cs b/src/Serval/src/Serval.Translation/Dtos/TranslationBuildDto.cs index 90a3a0b3..60976aab 100644 --- a/src/Serval/src/Serval.Translation/Dtos/TranslationBuildDto.cs +++ b/src/Serval/src/Serval.Translation/Dtos/TranslationBuildDto.cs @@ -33,6 +33,7 @@ public record TranslationBuildDto /// } /// public object? Options { get; init; } + public string? Model { get; init; } public string? DeploymentVersion { get; init; } public required ExecutionDataDto ExecutionData { get; init; } public IReadOnlyList? Phases { get; init; } diff --git a/src/Serval/src/Serval.Translation/Features/Engines/StartBuild.cs b/src/Serval/src/Serval.Translation/Features/Engines/StartBuild.cs index 275bd648..ce3f9062 100644 --- a/src/Serval/src/Serval.Translation/Features/Engines/StartBuild.cs +++ b/src/Serval/src/Serval.Translation/Features/Engines/StartBuild.cs @@ -43,6 +43,7 @@ public class StartBuildHandler( IEngineServiceFactory engineFactory, ILogger logger, DtoMapper dtoMapper, + IIdGenerator idGenerator, IConfiguration configuration ) : IRequestHandler { @@ -72,6 +73,7 @@ await builds.ExistsAsync( Build build = new() { + Id = idGenerator.GenerateId(), EngineRef = engine.Id, Owner = engine.Owner, Name = request.BuildConfig.Name, @@ -80,9 +82,7 @@ await builds.ExistsAsync( Options = MapOptions(request.BuildConfig.Options), DeploymentVersion = configuration.GetValue("deploymentVersion") ?? "Unknown", DateCreated = DateTime.UtcNow, - Model = request.BuildConfig.Model, }; - await builds.InsertAsync(build, ct); IReadOnlyList corpora = contractMapper.Map(build, engine); @@ -119,9 +119,13 @@ await builds.ExistsAsync( logger.LogInformation("Error parsing build request summary."); } - await engineFactory + StartBuildContract result = await engineFactory .GetEngineService(engine.Type) .StartBuildAsync(engine.Id, build.Id, corpora, buildOptions, build.Model, ct); + + build = build with { Model = result.Model }; + await builds.InsertAsync(build, ct); + return new StartBuildResponse(dtoMapper.Map(build)); }, cancellationToken diff --git a/src/Serval/src/Serval.Translation/Services/DtoMapper.cs b/src/Serval/src/Serval.Translation/Services/DtoMapper.cs index bfe7832f..08857f71 100644 --- a/src/Serval/src/Serval.Translation/Services/DtoMapper.cs +++ b/src/Serval/src/Serval.Translation/Services/DtoMapper.cs @@ -53,6 +53,7 @@ public TranslationBuildDto Map(Build source) DateCompleted = source.DateCompleted, DateFinished = source.DateFinished, Options = source.Options, + Model = source.Model, DeploymentVersion = source.DeploymentVersion, ExecutionData = Map(source.ExecutionData), Phases = source.Phases?.Select(Map).ToList(), diff --git a/src/Serval/src/Serval.Translation/Services/PlatformService.cs b/src/Serval/src/Serval.Translation/Services/PlatformService.cs index 6feb12d7..dad32125 100644 --- a/src/Serval/src/Serval.Translation/Services/PlatformService.cs +++ b/src/Serval/src/Serval.Translation/Services/PlatformService.cs @@ -485,7 +485,7 @@ public async Task InsertPretranslationsAsync( { { "bookId", b.bookId }, { "averagePretranslationConfidence", b.averageConfidence }, - { "modelName", model ?? "Unknown" }, + { "modelName", model ?? "unknown" }, } ) ) diff --git a/src/Serval/test/Serval.ApiServer.IntegrationTests/TranslationEngineTests.cs b/src/Serval/test/Serval.ApiServer.IntegrationTests/TranslationEngineTests.cs index a662db79..420db37b 100644 --- a/src/Serval/test/Serval.ApiServer.IntegrationTests/TranslationEngineTests.cs +++ b/src/Serval/test/Serval.ApiServer.IntegrationTests/TranslationEngineTests.cs @@ -2742,6 +2742,16 @@ public TestEnvironment() EchoService .GetLanguageInfoAsync(Arg.Any(), Arg.Any()) .Returns(Task.FromResult(languageInfo)); + EchoService + .StartBuildAsync( + Arg.Any(), + Arg.Any(), + Arg.Any>(), + Arg.Any(), + Arg.Any(), + Arg.Any() + ) + .Returns(Task.FromResult(new StartBuildContract())); NmtService = Substitute.For(); NmtService @@ -2757,11 +2767,31 @@ public TestEnvironment() NmtService .GetLanguageInfoAsync(Arg.Is("invalid_language"), Arg.Any()) .Returns(Task.FromException(new InvalidOperationException())); + NmtService + .StartBuildAsync( + Arg.Any(), + Arg.Any(), + Arg.Any>(), + Arg.Any(), + Arg.Any(), + Arg.Any() + ) + .Returns(Task.FromResult(new StartBuildContract { Model = "nllb" })); SmtService = Substitute.For(); SmtService .GetLanguageInfoAsync(Arg.Any(), Arg.Any()) .Returns(Task.FromResult(languageInfo)); + SmtService + .StartBuildAsync( + Arg.Any(), + Arg.Any(), + Arg.Any>(), + Arg.Any(), + Arg.Any(), + Arg.Any() + ) + .Returns(Task.FromResult(new StartBuildContract())); _dataFileOptions = _scope.ServiceProvider.GetRequiredService>(); ZipParatextProject(FILE3_FILENAME); ZipParatextProject(FILE4_FILENAME); diff --git a/src/Serval/test/Serval.E2ETests/ServalApiTests.cs b/src/Serval/test/Serval.E2ETests/ServalApiTests.cs index 62f3a3a6..a5379bfe 100644 --- a/src/Serval/test/Serval.E2ETests/ServalApiTests.cs +++ b/src/Serval/test/Serval.E2ETests/ServalApiTests.cs @@ -286,8 +286,9 @@ public async Task Nmt_Paratext(bool withAdditionalFiles) ], }, ]; + _helperClient.TranslationBuildConfig.Model = "nllb-600m"; _helperClient.TranslationBuildConfig.Options = - "{\"max_steps\":50, \"use_key_terms\":true, \"parent_model_name\": \"facebook/nllb-200-distilled-600M\", \"train_params\": {\"per_device_train_batch_size\":4}, \"generate_params\":{\"num_beams\": 2}}"; + "{\"max_steps\":50, \"use_key_terms\":true, \"train_params\": {\"per_device_train_batch_size\":4}, \"generate_params\":{\"num_beams\": 2}}"; string buildId = await _helperClient.BuildEngineAsync(engineId); TranslationBuild build = await _helperClient.TranslationEnginesClient.GetBuildAsync(engineId, buildId); diff --git a/src/Serval/test/Serval.Shared.Tests/Services/BuildDiagnosticServiceTests.cs b/src/Serval/test/Serval.Shared.Tests/Services/BuildDiagnosticServiceTests.cs index 8589d430..28a9da75 100644 --- a/src/Serval/test/Serval.Shared.Tests/Services/BuildDiagnosticServiceTests.cs +++ b/src/Serval/test/Serval.Shared.Tests/Services/BuildDiagnosticServiceTests.cs @@ -220,7 +220,7 @@ public void CreateDiagnostic_Model0003() var code = "MODEL-0003"; var data = new Dictionary { - { "averagePretranslationConfidence", 0.37 }, + { "averagePretranslationConfidence", 0.370111 }, { "bookId", "MAT" }, { "modelName", "test-model" }, }; @@ -229,11 +229,11 @@ public void CreateDiagnostic_Model0003() Assert.That(diagnostic.Code, Is.EqualTo(code)); Assert.That(diagnostic.Category, Is.EqualTo("MODEL")); - Assert.That(diagnostic.Severity, Is.EqualTo(Contracts.DiagnosticSeverity.Warn)); + Assert.That(diagnostic.Severity, Is.EqualTo(DiagnosticSeverity.Warn)); Assert.That( diagnostic.Message, Is.EqualTo( - "The average pretranslation model confidence 0.37 in book MAT is unusually low for the base model test-model." + "The average pretranslation model confidence 0.370 in book MAT is unusually low for the base model test-model." ) ); Assert.That(diagnostic.Data, Is.EqualTo(data)); diff --git a/src/Serval/test/Serval.Translation.Tests/Features/Engines/EnginesHandlersTests.cs b/src/Serval/test/Serval.Translation.Tests/Features/Engines/EnginesHandlersTests.cs index 9fbfdbdc..93dbf214 100644 --- a/src/Serval/test/Serval.Translation.Tests/Features/Engines/EnginesHandlersTests.cs +++ b/src/Serval/test/Serval.Translation.Tests/Features/Engines/EnginesHandlersTests.cs @@ -174,6 +174,7 @@ public async Task StartBuild_TrainOnNotSpecified() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -243,6 +244,7 @@ public async Task StartBuild_TextIdsEmpty() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -321,6 +323,7 @@ public async Task StartBuild_TextIdsPopulated() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -399,6 +402,7 @@ public async Task StartBuild_TextIdsNotSpecified() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -472,6 +476,7 @@ public async Task StartBuild_OneOfMultipleCorpora() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -549,6 +554,7 @@ public async Task StartBuild_TrainOnOnePretranslateTheOther() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -668,6 +674,7 @@ public async Task StartBuild_TextFilesScriptureRangeSpecified() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); Assert.ThrowsAsync(() => @@ -697,6 +704,7 @@ public async Task StartBuild_ScriptureRangeSpecified() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -775,6 +783,7 @@ public async Task StartBuild_ScriptureRangeEmptyString() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -853,6 +862,7 @@ public async Task StartBuild_ParallelCorpus_TextFiles() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -969,6 +979,7 @@ public async Task StartBuild_ParallelCorpus_OneOfMultipleCorpora() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1056,6 +1067,7 @@ public async Task StartBuild_ParallelCorpus_TrainOnOnePretranslateTheOther() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1185,6 +1197,7 @@ public async Task StartBuild_TextIds_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1301,6 +1314,7 @@ public async Task StartBuild_ScriptureRange_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1433,6 +1447,7 @@ public async Task StartBuild_MixedSourceAndTarget_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1557,6 +1572,7 @@ public async Task StartBuild_TextFilesScriptureRangeSpecified_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); Assert.ThrowsAsync(() => @@ -1610,6 +1626,7 @@ public async Task StartBuild_NoFilters_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1714,6 +1731,7 @@ public async Task StartBuild_TrainOnNotSpecified_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -1811,6 +1829,7 @@ public async Task StartBuild_NoTargetFilter_ParallelCorpus() env.EngineServiceFactory, Substitute.For>(), env.DtoMapper, + new ObjectIdGenerator(), Substitute.For() ); StartBuildResponse response = await handler.HandleAsync( @@ -2778,7 +2797,7 @@ public TestEnvironment() Arg.Any(), Arg.Any() ) - .Returns(Task.CompletedTask); + .Returns(Task.FromResult(new StartBuildContract())); TranslationEngineService .UpdateAsync(Arg.Any(), Arg.Any(), Arg.Any(), Arg.Any()) .Returns(Task.CompletedTask); diff --git a/src/Serval/test/Serval.Translation.Tests/Services/PlatformServiceTests.cs b/src/Serval/test/Serval.Translation.Tests/Services/PlatformServiceTests.cs index ce1435dd..cab9eb5a 100644 --- a/src/Serval/test/Serval.Translation.Tests/Services/PlatformServiceTests.cs +++ b/src/Serval/test/Serval.Translation.Tests/Services/PlatformServiceTests.cs @@ -97,7 +97,7 @@ await env.Builds.InsertAsync( Id = "b0", EngineRef = "e0", Owner = "owner1", - Model = "NLLB", + Model = "nllb", } ); @@ -123,7 +123,7 @@ await env.PlatformService.InsertPretranslationsAsync( executionData.Diagnostics[0].Data["averagePretranslationConfidence"], Is.EqualTo(0.2487).Within(0.0001) ); - Assert.That(executionData.Diagnostics[0].Data["modelName"], Is.EqualTo("NLLB")); + Assert.That(executionData.Diagnostics[0].Data["modelName"], Is.EqualTo("nllb")); } [Test]