mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-21 19:18:55 +00:00
### 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>
631 lines
29 KiB
C#
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
|
|
}
|
|
}
|