onnxruntime/csharp/test/Microsoft.ML.OnnxRuntime.Tests.Common/TrainingTest.cs
raoanag 424107a82a
Merge main to WindowsAI (#18122)
### Description
Merge main to WindowsAI



### Motivation and Context
<!-- - Why is this change required? What problem does it solve?
- If it fixes an open issue, please link to the issue here. -->

---------

Signed-off-by: Nash <george.nash@intel.com>
Signed-off-by: Yiming Hu <yiming.hu@amd.com>
Signed-off-by: Liqun Fu <liqfu@microsoft.com>
Co-authored-by: Kaz Nishimura <kazssym@linuxfront.com>
Co-authored-by: Tianlei Wu <tlwu@microsoft.com>
Co-authored-by: Nat Kershaw (MSFT) <nakersha@microsoft.com>
Co-authored-by: Yulong Wang <7679871+fs-eire@users.noreply.github.com>
Co-authored-by: Changming Sun <chasun@microsoft.com>
Co-authored-by: zesongw <zesong.wang@intel.com>
Co-authored-by: Yi Zhang <zhanyi@microsoft.com>
Co-authored-by: Dmitri Smirnov <yuslepukhin@users.noreply.github.com>
Co-authored-by: Yifan Li <109183385+yf711@users.noreply.github.com>
Co-authored-by: simonjub <78098752+simonjub@users.noreply.github.com>
Co-authored-by: PeixuanZuo <94887879+PeixuanZuo@users.noreply.github.com>
Co-authored-by: Adrian Lizarraga <adlizarraga@microsoft.com>
Co-authored-by: Edward Chen <18449977+edgchen1@users.noreply.github.com>
Co-authored-by: Arthur Islamov <arthur@islamov.ai>
Co-authored-by: Jambay Kinley <jambaykinley@microsoft.com>
Co-authored-by: Justin Chu <justinchuby@users.noreply.github.com>
Co-authored-by: Wei-Sheng Chin <wschin@outlook.com>
Co-authored-by: Bowen Bao <bowbao@microsoft.com>
Co-authored-by: Hariharan Seshadri <shariharan91@gmail.com>
Co-authored-by: Numfor Tiapo <numsmt2@gmail.com>
Co-authored-by: Vincent Wang <wangwchpku@outlook.com>
Co-authored-by: Pranav Sharma <prs@microsoft.com>
Co-authored-by: George Nash <george.nash@intel.com>
Co-authored-by: Abhishek Jindal <abjindal@microsoft.com>
Co-authored-by: pengwa <pengwa@microsoft.com>
Co-authored-by: Yiming Hu <woinck@users.noreply.github.com>
Co-authored-by: Jiajia Qin <jiajia.qin@intel.com>
Co-authored-by: Lukas Berbuer <36054362+lukasberbuer@users.noreply.github.com>
Co-authored-by: Wanming Lin <wanming.lin@intel.com>
Co-authored-by: Xavier Dupré <xadupre@users.noreply.github.com>
Co-authored-by: aimilefth <60664743+aimilefth@users.noreply.github.com>
Co-authored-by: Baiju Meswani <bmeswani@microsoft.com>
Co-authored-by: Adam Pocock <adam.pocock@oracle.com>
Co-authored-by: Chi Lo <54722500+chilo-ms@users.noreply.github.com>
Co-authored-by: RandySheriffH <48490400+RandySheriffH@users.noreply.github.com>
Co-authored-by: Randy Shuai <rashuai@microsoft.com>
Co-authored-by: Vadym Stupakov <vadim.stupakov@gmail.com>
Co-authored-by: Jian Chen <cjian@microsoft.com>
Co-authored-by: Brian Lambert <98757707+brian-pieces@users.noreply.github.com>
Co-authored-by: Nicolò Lucchesi <nicolo.lucchesi@gmail.com>
Co-authored-by: liqun Fu <liqfu@microsoft.com>
Co-authored-by: trajep <trajepl@gmail.com>
Co-authored-by: Scott McKay <skottmckay@gmail.com>
Co-authored-by: Mustafa Ateş Uzun <mustafauzun0@gmail.com>
Co-authored-by: MistEO <mistereo@hotmail.com>
Co-authored-by: satyajandhyala <satya.k.jandhyala@gmail.com>
Co-authored-by: shaahji <96227573+shaahji@users.noreply.github.com>
Co-authored-by: Rachel Guo <35738743+YUNQIUGUO@users.noreply.github.com>
Co-authored-by: rachguo <rachguo@rachguos-Mini.attlocal.net>
Co-authored-by: Caroline Zhu <wolfivyaura@gmail.com>
Co-authored-by: Caroline Zhu <carolinezhu@microsoft.com>
Co-authored-by: Guenther Schmuelling <guschmue@microsoft.com>
Co-authored-by: xhcao <xinghua.cao@intel.com>
Co-authored-by: Ella Charlaix <80481427+echarlaix@users.noreply.github.com>
Co-authored-by: Xu Xing <xing.xu@intel.com>
Co-authored-by: Hector Li <hecli@microsoft.com>
Co-authored-by: Ye Wang <52801275+wangyems@users.noreply.github.com>
Co-authored-by: Your Name <you@example.com>
Co-authored-by: Benedikt Hilmes <benedikt.hilmes@rwth-aachen.de>
Co-authored-by: rachguo <rachguo@rachguos-Mac-mini.local>
Co-authored-by: George Wu <jywu@microsoft.com>
Co-authored-by: JiCheng <wejoncy@163.com>
Co-authored-by: Sheil Kumar <smk2007@gmail.com>
Co-authored-by: Sheil Kumar <sheilk@microsoft.com>
Co-authored-by: cloudhan <guangyunhan@microsoft.com>
Co-authored-by: kyoshisuki <143475866+kyoshisuki@users.noreply.github.com>
Co-authored-by: aciddelgado <139922440+aciddelgado@users.noreply.github.com>
Co-authored-by: tlwu@microsoft.com <tlwu@a100.crj0ad2y1kku1j4yxl4sj10o4e.gx.internal.cloudapp.net>
Co-authored-by: Maximilian Müller <44298237+gedoensmax@users.noreply.github.com>
Co-authored-by: Tang, Cheng <souptc@gmail.com>
Co-authored-by: Cheng Tang <chenta@microsoft.com@orttrainingdev9.d32nl1ml4oruzj4qz3bqlggovf.px.internal.cloudapp.net>
Co-authored-by: Cheng Tang <chenta@microsoft.com>
Co-authored-by: Jeff Daily <jeff.daily@amd.com>
Co-authored-by: cloudhan <cloudhan@outlook.com>
Co-authored-by: Yufeng Li <liyufeng1987@gmail.com>
Co-authored-by: Zhang Lei <zhang.huanning@hotmail.com>
Co-authored-by: Dwayne Robinson <fdwr@hotmail.com>
Co-authored-by: Zhipeng Han <zhipeng.han@outlook.com>
Co-authored-by: Thiago Crepaldi <thiago.crepaldi@microsoft.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Patrice Vignola <vignola.patrice@gmail.com>
Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com>
Co-authored-by: snadampal <87143774+snadampal@users.noreply.github.com>
Co-authored-by: Sumit Agarwal <sumitagarwal330@gmail.com>
Co-authored-by: Ashwini Khade <askhade@microsoft.com>
Co-authored-by: Yang Gu <yang.gu@intel.com>
Co-authored-by: Cheng Tang <chenta@a100.crj0ad2y1kku1j4yxl4sj10o4e.gx.internal.cloudapp.net>
Co-authored-by: mindest <30493312+mindest@users.noreply.github.com>
Co-authored-by: Scott McKay <Scott.McKay@microsoft.com>
Co-authored-by: Xavier Dupre <xadupre@microsoft.com@orttrainingdev9.d32nl1ml4oruzj4qz3bqlggovf.px.internal.cloudapp.net>
Co-authored-by: guyang3532 <62738430+guyang3532@users.noreply.github.com>
Co-authored-by: Carson M <carson@pyke.io>
Co-authored-by: sophies927 <107952697+sophies927@users.noreply.github.com>
2023-10-27 17:08:01 -07:00

631 lines
29 KiB
C#

// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
using System;
using System.IO;
using Xunit;
using Xunit.Abstractions;
#if __TRAINING_ENABLED_NATIVE_BUILD__
using Microsoft.ML.OnnxRuntime.Tensors;
using System.Collections.Generic;
using System.Linq;
#endif
// This runs in a separate package built from EndToEndTests
// and for this reason it can not refer to non-public members
// of Onnxruntime package
namespace Microsoft.ML.OnnxRuntime.Tests
{
public partial class TrainingTest
{
private readonly ITestOutputHelper output;
public TrainingTest(ITestOutputHelper o)
{
this.output = o;
}
#if !__TRAINING_ENABLED_NATIVE_BUILD__
[Fact(DisplayName = "TestLoadCheckpointThrows")]
public void TestLoadCheckpointThrows()
{
string path = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
var ex = Assert.Throws<InvalidOperationException>(() => { var opt = CheckpointState.LoadCheckpoint(path); });
Assert.Contains("Please install the Microsoft.ML.OnnxRuntime.Training NuGet package.", ex.Message);
}
#endif
#if __TRAINING_ENABLED_NATIVE_BUILD__
[Fact(DisplayName = "TestLoadCheckpoint")]
public void TestLoadCheckpoint()
{
string path = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var opt = CheckpointState.LoadCheckpoint(path))
{
Assert.NotNull(opt);
}
}
[Fact(DisplayName = "TestCreateTrainingSession")]
public void TestCreateTrainingSession()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
}
}
[Fact(DisplayName = "TestTrainingSessionTrainStep")]
public void TestTrainingSessionTrainStep()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
float[] expectedOutput = TestDataLoader.LoadTensorFromFile("loss_1.out");
var expectedOutputDimensions = new int[] { 1 };
float[] input = TestDataLoader.LoadTensorFromFile("input-0.in");
Int32[] labels = { 1, 1 };
// Run inference with pinned inputs and pinned outputs
using (DisposableListTest<FixedBufferOnnxValue> pinnedInputs = new DisposableListTest<FixedBufferOnnxValue>(),
pinnedOutputs = new DisposableListTest<FixedBufferOnnxValue>())
{
var memInfo = OrtMemoryInfo.DefaultInstance; // CPU
// Create inputs
long[] inputShape = { 2, 784 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<float>(memInfo, input,
TensorElementType.Float, inputShape, input.Length * sizeof(float)));
long[] labelsShape = { 2 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<Int32>(memInfo, labels,
TensorElementType.Int32, labelsShape, labels.Length * sizeof(Int32)));
// Prepare output buffer
long[] outputShape = { };
float[] outputBuffer = new float[expectedOutput.Length];
pinnedOutputs.Add(FixedBufferOnnxValue.CreateFromMemory<float>(memInfo, outputBuffer,
TensorElementType.Float, outputShape, outputBuffer.Length * sizeof(float)));
trainingSession.TrainStep(pinnedInputs, pinnedOutputs);
Assert.Equal(expectedOutput, outputBuffer, new FloatComparer());
}
}
}
void RunTrainStep(TrainingSession trainingSession)
{
float[] expectedOutput = TestDataLoader.LoadTensorFromFile("loss_1.out");
var expectedOutputDimensions = new int[] { 1 };
float[] input = TestDataLoader.LoadTensorFromFile("input-0.in");
Int32[] labels = { 1, 1 };
// Run inference with pinned inputs and pinned outputs
using (DisposableListTest<FixedBufferOnnxValue> pinnedInputs = new DisposableListTest<FixedBufferOnnxValue>())
{
var memInfo = OrtMemoryInfo.DefaultInstance; // CPU
// Create inputs
long[] inputShape = { 2, 784 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<float>(memInfo, input,
TensorElementType.Float, inputShape, input.Length * sizeof(float)));
long[] labelsShape = { 2 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<Int32>(memInfo, labels,
TensorElementType.Int32, labelsShape, labels.Length * sizeof(Int32)));
var outputs = trainingSession.TrainStep(pinnedInputs);
trainingSession.LazyResetGrad();
outputs = trainingSession.TrainStep(pinnedInputs);
var outputBuffer = outputs.ElementAtOrDefault(0);
Assert.Equal("onnx::loss::21273", outputBuffer.Name);
Assert.Equal(OnnxValueType.ONNX_TYPE_TENSOR, outputBuffer.ValueType);
Assert.Equal(TensorElementType.Float, outputBuffer.ElementType);
var outLabelTensor = outputBuffer.AsTensor<float>();
Assert.NotNull(outLabelTensor);
Assert.Equal(expectedOutput, outLabelTensor.ToArray(), new FloatComparer());
}
}
[Fact(DisplayName = "TestTrainingSessionTrainStepOrtOutput")]
public void TestTrainingSessionTrainStepOrtOutput()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
RunTrainStep(trainingSession);
}
}
[Fact(DisplayName = "TestSaveCheckpoint")]
public void TestSaveCheckpoint()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
// Save checkpoint
string savedCheckpointPath = Path.Combine(Directory.GetCurrentDirectory(), "saved_checkpoint.ckpt");
CheckpointState.SaveCheckpoint(state, savedCheckpointPath, true);
// Load checkpoint and run train step
var loadedState = CheckpointState.LoadCheckpoint(savedCheckpointPath);
cleanUp.Add(loadedState);
var newTrainingSession = new TrainingSession(loadedState, trainingPath, optimizerPath);
cleanUp.Add(newTrainingSession);
RunTrainStep(newTrainingSession);
}
}
[Fact(DisplayName = "TestTrainingSessionOptimizerStep")]
public void TestTrainingSessionOptimizerStep()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
float[] expectedOutput_1 = TestDataLoader.LoadTensorFromFile("loss_1.out");
float[] expectedOutput_2 = TestDataLoader.LoadTensorFromFile("loss_2.out");
var expectedOutputDimensions = new int[] { 1 };
float[] input = TestDataLoader.LoadTensorFromFile("input-0.in");
Int32[] labels = { 1, 1 };
// Run train step with pinned inputs and pinned outputs
using (DisposableListTest<FixedBufferOnnxValue> pinnedInputs = new DisposableListTest<FixedBufferOnnxValue>(),
pinnedOutputs = new DisposableListTest<FixedBufferOnnxValue>())
{
var memInfo = OrtMemoryInfo.DefaultInstance; // CPU
// Create inputs
long[] inputShape = { 2, 784 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<float>(memInfo, input,
TensorElementType.Float, inputShape, input.Length * sizeof(float)));
long[] labelsShape = { 2 };
pinnedInputs.Add(FixedBufferOnnxValue.CreateFromMemory<Int32>(memInfo, labels,
TensorElementType.Int32, labelsShape, labels.Length * sizeof(Int32)));
// Prepare output buffer
long[] outputShape = { };
float[] outputBuffer = new float[expectedOutput_1.Length];
pinnedOutputs.Add(FixedBufferOnnxValue.CreateFromMemory<float>(memInfo, outputBuffer,
TensorElementType.Float, outputShape, outputBuffer.Length * sizeof(float)));
trainingSession.TrainStep(pinnedInputs, pinnedOutputs);
Assert.Equal(expectedOutput_1, outputBuffer, new FloatComparer());
trainingSession.LazyResetGrad();
trainingSession.TrainStep(pinnedInputs, pinnedOutputs);
Assert.Equal(expectedOutput_1, outputBuffer, new FloatComparer());
trainingSession.OptimizerStep();
trainingSession.TrainStep(pinnedInputs, pinnedOutputs);
Assert.Equal(expectedOutput_2, outputBuffer, new FloatComparer());
}
}
}
[Fact(DisplayName = "TestTrainingSessionSetLearningRate")]
public void TestTrainingSessionSetLearningRate()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
float learningRate = 0.245f;
trainingSession.SetLearningRate(learningRate);
var actualLearningRate = trainingSession.GetLearningRate();
Assert.Equal(learningRate, actualLearningRate);
}
}
[Fact(DisplayName = "TestTrainingSessionLinearLRScheduler")]
public void TestTrainingSessionLinearLRScheduler()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
float learningRate = 0.1f;
trainingSession.RegisterLinearLRScheduler(2, 4, learningRate);
RunTrainStep(trainingSession);
trainingSession.OptimizerStep();
trainingSession.SchedulerStep();
Assert.Equal(0.05f, trainingSession.GetLearningRate());
trainingSession.OptimizerStep();
trainingSession.SchedulerStep();
Assert.Equal(0.1f, trainingSession.GetLearningRate());
trainingSession.OptimizerStep();
trainingSession.SchedulerStep();
Assert.Equal(0.05f, trainingSession.GetLearningRate());
trainingSession.OptimizerStep();
trainingSession.SchedulerStep();
Assert.Equal(0.0f, trainingSession.GetLearningRate());
}
}
[Fact(DisplayName = "TestTrainingSessionExportModelForInferencing")]
public void TestTrainingSessionExportModelForInferencing()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath);
cleanUp.Add(trainingSession);
var graphOutputs = new List<string>(){"output-0"};
string inferencePath = Path.Combine(Directory.GetCurrentDirectory(), "inference_model.onnx");
trainingSession.ExportModelForInferencing(inferencePath, graphOutputs);
Assert.True(File.Exists(inferencePath));
}
}
[Fact(DisplayName = "TestCheckpointStateAddProperty")]
public void TestCheckpointStateAddProperty()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string propertyName = "days in a week";
state.AddProperty(propertyName, (long)7);
var value = state.GetProperty(propertyName);
Assert.True(value is long);
Assert.Equal((long)7, value);
}
}
[Fact(DisplayName = "TestCheckpointStateAddFloatProperty")]
public void TestCheckpointStateAddFloatProperty()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string propertyName = "pi";
state.AddProperty(propertyName, (float)3.14);
var value = state.GetProperty(propertyName);
Assert.True(value is float);
Assert.Equal((float)3.14, value);
}
}
[Fact(DisplayName = "TestCheckpointStateAddStringProperty")]
public void TestCheckpointStateAddStringProperty()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string propertyName = "best ai framework";
state.AddProperty(propertyName, "onnxruntime");
var value = state.GetProperty(propertyName);
Assert.True(value is string);
Assert.Equal("onnxruntime", value);
}
}
[Fact(DisplayName = "TestTrainModelInputNames")]
public void TestTrainModelInputNames()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
var inputNames = trainingSession.InputNames(true);
Assert.True(inputNames.Count == 2);
Assert.Equal("input-0", inputNames[0]);
Assert.Equal("labels", inputNames[1]);
}
}
[Fact(DisplayName = "TestEvalModelInputNames")]
public void TestEvalModelInputNames()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath);
cleanUp.Add(trainingSession);
var inputNames = trainingSession.InputNames(false);
Assert.True(inputNames.Count == 2);
Assert.Equal("input-0", inputNames[0]);
Assert.Equal("labels", inputNames[1]);
}
}
[Fact(DisplayName = "TestTrainModelOutputNames")]
public void TestTrainModelOutputNames()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, optimizerPath);
cleanUp.Add(trainingSession);
var outputNames = trainingSession.OutputNames(true);
Assert.Single(outputNames);
Assert.Equal("onnx::loss::21273", outputNames[0]);
}
}
[Fact(DisplayName = "TestEvalModelOutputNames")]
public void TestEvalModelOutputNames()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
using (var cleanUp = new DisposableListTest<IDisposable>())
{
var state = CheckpointState.LoadCheckpoint(checkpointPath);
cleanUp.Add(state);
Assert.NotNull(state);
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath);
cleanUp.Add(trainingSession);
var outputNames = trainingSession.OutputNames(false);
Assert.Single(outputNames);
Assert.Equal("onnx::loss::21273", outputNames[0]);
}
}
[Fact(DisplayName = "TestToBuffer")]
public void TestToBuffer()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
using (var state = CheckpointState.LoadCheckpoint(checkpointPath))
using (var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath))
{
Assert.NotNull(state);
using (var buffer = trainingSession.ToBuffer(true))
{
Assert.NotNull(buffer);
var typeShape = buffer.GetTensorTypeAndShape();
Assert.Equal(1, typeShape.DimensionsCount);
var fetchedShape = typeShape.Shape;
Assert.Equal(397510, fetchedShape[0]);
}
}
}
[Fact(DisplayName = "TestFromBuffer")]
public void TestFromBuffer()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
using (var state = CheckpointState.LoadCheckpoint(checkpointPath))
using (var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath))
{
Assert.NotNull(state);
using (var buffer = trainingSession.ToBuffer(true))
{
Assert.NotNull(buffer);
var typeShape = buffer.GetTensorTypeAndShape();
Assert.Equal(1, typeShape.DimensionsCount);
var fetchedShape = typeShape.Shape;
Assert.Equal(397510, fetchedShape[0]);
trainingSession.FromBuffer(buffer, true);
}
}
}
[Fact(DisplayName = "TestSetSeed")]
public void TestSetSeed()
{
TrainingUtils.SetSeed(8888);
}
[Fact(DisplayName = "TestGetParameter")]
public void TestGetParameter()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
using (var state = CheckpointState.LoadCheckpoint(checkpointPath))
using (var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath))
using (var parameter = state.GetParameter("fc1.weight"))
{
Assert.NotNull(state);
Assert.NotNull(parameter);
var typeShape = parameter.GetTensorTypeAndShape();
Assert.Equal(2, typeShape.DimensionsCount);
var fetchedShape = typeShape.Shape;
Assert.Equal(500, fetchedShape[0]);
Assert.Equal(784, fetchedShape[1]);
}
}
[Fact(DisplayName = "TestUpdateParameter")]
public void TestUpdateParameter()
{
string checkpointPath = Path.Combine(Directory.GetCurrentDirectory(), "checkpoint.ckpt");
string trainingPath = Path.Combine(Directory.GetCurrentDirectory(), "training_model.onnx");
string evalPath = Path.Combine(Directory.GetCurrentDirectory(), "eval_model.onnx");
string optimizerPath = Path.Combine(Directory.GetCurrentDirectory(), "adamw.onnx");
using (var state = CheckpointState.LoadCheckpoint(checkpointPath))
using (var trainingSession = new TrainingSession(state, trainingPath, evalPath, optimizerPath))
{
Assert.NotNull(state);
using (var parameter = state.GetParameter("fc1.weight"))
{
Assert.NotNull(parameter);
var typeShape = parameter.GetTensorTypeAndShape();
Assert.Equal(2, typeShape.DimensionsCount);
var fetchedShape = typeShape.Shape;
Assert.Equal(500, fetchedShape[0]);
Assert.Equal(784, fetchedShape[1]);
float maxVal = 20;
Random randNum = new Random();
float[] updated_parameter_buffer = Enumerable
.Repeat(0, 500 * 784)
.Select(i => maxVal * (float)randNum.NextDouble())
.ToArray();
using (var updated_parameter = OrtValue.CreateTensorValueFromMemory(updated_parameter_buffer, fetchedShape))
{
state.UpdateParameter("fc1.weight", updated_parameter);
using (var current_parameter = state.GetParameter("fc1.weight"))
{
var current_parameter_tensor = current_parameter.GetTensorDataAsSpan<float>().ToArray();
Assert.Equal(updated_parameter_buffer, current_parameter_tensor);
Assert.NotEqual(parameter.GetTensorDataAsSpan<float>().ToArray(), current_parameter_tensor);
}
state.UpdateParameter("fc1.weight", parameter);
using (var current_parameter = state.GetParameter("fc1.weight"))
{
var current_parameter_tensor = current_parameter.GetTensorDataAsSpan<float>().ToArray();
Assert.Equal(parameter.GetTensorDataAsSpan<float>().ToArray(), current_parameter_tensor);
Assert.NotEqual(updated_parameter_buffer, current_parameter_tensor);
}
}
}
}
}
internal class FloatComparer : IEqualityComparer<float>
{
private float atol = 1e-3f;
private float rtol = 1.7e-2f;
public bool Equals(float x, float y)
{
return Math.Abs(x - y) <= (atol + rtol * Math.Abs(y));
}
public int GetHashCode(float x)
{
return x.GetHashCode();
}
}
#endif
}
}