From 882c4f79121e5cdac372ff65e1c7505666830b8d Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Wed, 19 Sep 2018 10:11:51 -0700 Subject: [PATCH 1/7] Some changes for adding experimental conversion to ONNX 1. Introduce a new argument to SaveOnnx, which is OnnxVersion. Two values are currently allowed, "Latest" and "Experimental". Note that "Latest" means that the produced ONNX model meets the latest ONNX release while "Experimental" may produce things not officially supported in ONNX. 2. For (1), the interface of saving ONNX is slightly changed. Now, CanSaveOnnx requires an OnnxContext as its input argument. 3. Add exporter for ConvertTransform. It doesn't use standard ONNX operator. --- .../DataView/RowToRowMapperTransform.cs | 4 +-- .../Model/Onnx/ICanSaveOnnx.cs | 2 +- .../Model/Onnx/OnnxContext.cs | 8 ++++++ .../Prediction/Calibrator.cs | 10 +++---- .../Scorers/GenericScorer.cs | 2 +- .../Scorers/MultiClassClassifierScorer.cs | 4 +-- .../Scorers/PredictedLabelScorerBase.cs | 2 +- .../Scorers/SchemaBindablePredictorWrapper.cs | 2 +- .../Transforms/ConcatTransform.cs | 4 +-- .../Transforms/ConvertTransform.cs | 25 ++++++++++++++++++ .../Transforms/KeyToVectorTransform.cs | 2 +- .../Transforms/NormalizeColumn.cs | 6 ++--- .../Transforms/NormalizeUtils.cs | 2 +- .../Transforms/Normalizer.cs | 6 ++--- .../Transforms/TermTransform.cs | 2 +- .../Transforms/TransformBase.cs | 4 +-- src/Microsoft.ML.FastTree/FastTree.cs | 2 +- src/Microsoft.ML.Onnx/OnnxContextImpl.cs | 6 ++++- src/Microsoft.ML.Onnx/OnnxUtils.cs | 26 ++++++++++++++++++- src/Microsoft.ML.Onnx/SaveOnnxCommand.cs | 24 ++++++++++------- .../Standard/LinearPredictor.cs | 2 +- .../MulticlassLogisticRegression.cs | 2 +- .../NAReplaceTransform.cs | 2 +- 23 files changed, 107 insertions(+), 42 deletions(-) diff --git a/src/Microsoft.ML.Data/DataView/RowToRowMapperTransform.cs b/src/Microsoft.ML.Data/DataView/RowToRowMapperTransform.cs index ef19012bf6..13240a626d 100644 --- a/src/Microsoft.ML.Data/DataView/RowToRowMapperTransform.cs +++ b/src/Microsoft.ML.Data/DataView/RowToRowMapperTransform.cs @@ -246,7 +246,7 @@ private static VersionInfo GetVersionInfo() public override ISchema Schema { get { return _bindings; } } - public bool CanSaveOnnx => _mapper is ICanSaveOnnx onnxMapper ? onnxMapper.CanSaveOnnx : false; + public bool CanSaveOnnx(OnnxContext ctx) => _mapper is ICanSaveOnnx onnxMapper ? onnxMapper.CanSaveOnnx(ctx) : false; public bool CanSavePfa => _mapper is ICanSavePfa pfaMapper ? pfaMapper.CanSavePfa : false; @@ -338,7 +338,7 @@ public void SaveAsOnnx(OnnxContext ctx) Host.CheckValue(ctx, nameof(ctx)); if (_mapper is ISaveAsOnnx onnx) { - Host.Check(onnx.CanSaveOnnx, "Cannot be saved as ONNX."); + Host.Check(onnx.CanSaveOnnx(ctx), "Cannot be saved as ONNX."); onnx.SaveAsOnnx(ctx); } } diff --git a/src/Microsoft.ML.Data/Model/Onnx/ICanSaveOnnx.cs b/src/Microsoft.ML.Data/Model/Onnx/ICanSaveOnnx.cs index 103f2efd9f..f5163ac022 100644 --- a/src/Microsoft.ML.Data/Model/Onnx/ICanSaveOnnx.cs +++ b/src/Microsoft.ML.Data/Model/Onnx/ICanSaveOnnx.cs @@ -15,7 +15,7 @@ public interface ICanSaveOnnx /// only detectable during runtime that would prevent its being savable. (E.g., /// it may wrap some other object that may or may not be savable.) /// - bool CanSaveOnnx { get; } + bool CanSaveOnnx(OnnxContext ctx); } /// diff --git a/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs b/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs index 230f2600a3..9dc62ac55b 100644 --- a/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs +++ b/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs @@ -7,6 +7,8 @@ namespace Microsoft.ML.Runtime.Model.Onnx { + public enum OnnxVersion { Latest, Experimental } + /// /// A context for defining a ONNX output. The context internally contains the model-in-progress being built. This /// same context object is iteratively given to exportable components via the interface @@ -98,5 +100,11 @@ public abstract OnnxNode CreateNode(string opType, IEnumerable inputs, /// A node added to the in-progress ONNX graph, that attributes can be set on public OnnxNode CreateNode(string opType, string input, string output, string name, string domain = null) => CreateNode(opType, new[] { input }, new[] { output }, name, domain); + + /// + /// Get the targeted ONNX version string. Only two values are allowed now: "latest" and "experimental". + /// + /// + public abstract OnnxVersion GetOnnxVersion(); } } diff --git a/src/Microsoft.ML.Data/Prediction/Calibrator.cs b/src/Microsoft.ML.Data/Prediction/Calibrator.cs index 6cee31b9d2..24fe32d2c7 100644 --- a/src/Microsoft.ML.Data/Prediction/Calibrator.cs +++ b/src/Microsoft.ML.Data/Prediction/Calibrator.cs @@ -225,7 +225,7 @@ public abstract class ValueMapperCalibratedPredictorBase : CalibratedPredictorBa public ColumnType OutputType => _mapper.OutputType; public ColumnType DistType => NumberType.Float; public bool CanSavePfa => (_mapper as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (_mapper as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (_mapper as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; protected ValueMapperCalibratedPredictorBase(IHostEnvironment env, string name, IPredictorProducing predictor, ICalibrator calibrator) : base(env, name, predictor, calibrator) @@ -308,7 +308,7 @@ public bool SaveAsOnnx(OnnxContext ctx, string[] outputNames, string featureColu return false; var calibrator = Calibrator as ISingleCanSaveOnnx; - if (!(calibrator?.CanSaveOnnx == true && calibrator.SaveAsOnnx(ctx, new[] { outputNames[1], outputNames[2] }, featureColumnName))) + if (!(calibrator?.CanSaveOnnx(ctx) == true && calibrator.SaveAsOnnx(ctx, new[] { outputNames[1], outputNames[2] }, featureColumnName))) ctx.RemoveVariable(outputNames[1], true); return true; @@ -617,7 +617,7 @@ private static VersionInfo GetVersionInfo() /// public bool CanSavePfa => (_bindable as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (_bindable as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (_bindable as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; public SchemaBindableCalibratedPredictor(IHostEnvironment env, IPredictorProducing predictor, ICalibrator calibrator) : base(env, LoaderSignature, predictor, calibrator) @@ -663,7 +663,7 @@ public bool SaveAsOnnx(OnnxContext ctx, RoleMappedSchema schema, string[] output Host.CheckValue(ctx, nameof(ctx)); Host.CheckParam(Utils.Size(outputs) == 2, nameof(outputs), "Expected this to have two outputs"); Host.CheckValue(schema, nameof(schema)); - Host.Check(CanSaveOnnx, "Called despite not being savable"); + Host.Check(CanSaveOnnx(ctx), "Called despite not being savable"); return false; } @@ -1342,7 +1342,7 @@ private static VersionInfo GetVersionInfo() public Double ParamA { get; } public Double ParamB { get; } public bool CanSavePfa => true; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public PlattCalibrator(IHostEnvironment env, Double paramA, Double paramB) { diff --git a/src/Microsoft.ML.Data/Scorers/GenericScorer.cs b/src/Microsoft.ML.Data/Scorers/GenericScorer.cs index 41c12e94ed..39cbc746b1 100644 --- a/src/Microsoft.ML.Data/Scorers/GenericScorer.cs +++ b/src/Microsoft.ML.Data/Scorers/GenericScorer.cs @@ -140,7 +140,7 @@ private static VersionInfo GetVersionInfo() public bool CanSavePfa => (Bindable as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (Bindable as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (Bindable as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; /// /// The entry point for creating a . diff --git a/src/Microsoft.ML.Data/Scorers/MultiClassClassifierScorer.cs b/src/Microsoft.ML.Data/Scorers/MultiClassClassifierScorer.cs index c12fd9b4d1..28f6fc160c 100644 --- a/src/Microsoft.ML.Data/Scorers/MultiClassClassifierScorer.cs +++ b/src/Microsoft.ML.Data/Scorers/MultiClassClassifierScorer.cs @@ -77,7 +77,7 @@ public sealed class LabelNameBindableMapper : ISchemaBindableMapper, ICanSaveMod public VectorType Type => _type; public bool CanSavePfa => (_bindable as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (_bindable as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (_bindable as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; public ISchemaBindableMapper InnerBindable => _bindable; private static VersionInfo GetVersionInfo() @@ -207,7 +207,7 @@ public bool SaveAsOnnx(OnnxContext ctx, RoleMappedSchema schema, string[] output { Contracts.CheckValue(ctx, nameof(ctx)); Contracts.CheckValue(schema, nameof(schema)); - Contracts.Check(CanSaveOnnx, "Cannot be saved as ONNX."); + Contracts.Check(CanSaveOnnx(ctx), "Cannot be saved as ONNX."); Contracts.Assert(_bindable is IBindableCanSaveOnnx); return ((IBindableCanSaveOnnx)_bindable).SaveAsOnnx(ctx, schema, outputNames); } diff --git a/src/Microsoft.ML.Data/Scorers/PredictedLabelScorerBase.cs b/src/Microsoft.ML.Data/Scorers/PredictedLabelScorerBase.cs index 2fd039897a..95e072f39f 100644 --- a/src/Microsoft.ML.Data/Scorers/PredictedLabelScorerBase.cs +++ b/src/Microsoft.ML.Data/Scorers/PredictedLabelScorerBase.cs @@ -284,7 +284,7 @@ protected override BindingsBase GetBindings() public bool CanSavePfa => (Bindable as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (Bindable as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (Bindable as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; protected PredictedLabelScorerBase(ScorerArgumentsBase args, IHostEnvironment env, IDataView data, ISchemaBoundMapper mapper, RoleMappedSchema trainSchema, string registrationName, string scoreColKind, string scoreColName, diff --git a/src/Microsoft.ML.Data/Scorers/SchemaBindablePredictorWrapper.cs b/src/Microsoft.ML.Data/Scorers/SchemaBindablePredictorWrapper.cs index 1e07f587f7..0a04989a4e 100644 --- a/src/Microsoft.ML.Data/Scorers/SchemaBindablePredictorWrapper.cs +++ b/src/Microsoft.ML.Data/Scorers/SchemaBindablePredictorWrapper.cs @@ -46,7 +46,7 @@ public abstract class SchemaBindablePredictorWrapperBase : ISchemaBindableMapper public bool CanSavePfa => (ValueMapper as ICanSavePfa)?.CanSavePfa == true; - public bool CanSaveOnnx => (ValueMapper as ICanSaveOnnx)?.CanSaveOnnx == true; + public bool CanSaveOnnx(OnnxContext ctx) => (ValueMapper as ICanSaveOnnx)?.CanSaveOnnx(ctx) == true; public SchemaBindablePredictorWrapperBase(IPredictor predictor) { diff --git a/src/Microsoft.ML.Data/Transforms/ConcatTransform.cs b/src/Microsoft.ML.Data/Transforms/ConcatTransform.cs index e0f6c1dfdf..e152106af1 100644 --- a/src/Microsoft.ML.Data/Transforms/ConcatTransform.cs +++ b/src/Microsoft.ML.Data/Transforms/ConcatTransform.cs @@ -424,7 +424,7 @@ private sealed class Mapper : IRowMapper, ISaveAsOnnx, ISaveAsPfa private readonly ConcatTransform _parent; private readonly BoundColumn[] _columns; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public bool CanSavePfa => true; public Mapper(ConcatTransform parent, ISchema inputSchema) @@ -895,7 +895,7 @@ public void SaveAsPfa(BoundPfaContext ctx) public void SaveAsOnnx(OnnxContext ctx) { _host.CheckValue(ctx, nameof(ctx)); - Contracts.Assert(CanSaveOnnx); + Contracts.Assert(CanSaveOnnx(ctx)); string opType = "FeatureVectorizer"; for (int iinfo = 0; iinfo < _columns.Length; ++iinfo) diff --git a/src/Microsoft.ML.Data/Transforms/ConvertTransform.cs b/src/Microsoft.ML.Data/Transforms/ConvertTransform.cs index 52005c7558..b1047ef67b 100644 --- a/src/Microsoft.ML.Data/Transforms/ConvertTransform.cs +++ b/src/Microsoft.ML.Data/Transforms/ConvertTransform.cs @@ -16,6 +16,7 @@ using Microsoft.ML.Runtime.Data.Conversion; using Microsoft.ML.Runtime.Internal.Utilities; using Microsoft.ML.Runtime.Model; +using Microsoft.ML.Runtime.Model.Onnx; using Microsoft.ML.Runtime.Command; using Microsoft.ML.Runtime.EntryPoints; @@ -374,6 +375,30 @@ public override void Save(ModelSaveContext ctx) } } + public override bool CanSaveOnnx(OnnxContext ctx) => ctx.GetOnnxVersion() == OnnxVersion.Experimental; + + protected override bool SaveAsOnnxCore(OnnxContext ctx, int iinfo, ColInfo info, string srcVariableName, string dstVariableName) + { + var opType = "CSharp"; + + for (int i = 0; i < _exes.Length; i++) + { + var ex = _exes[i]; + var node = ctx.CreateNode(opType, srcVariableName, dstVariableName, ctx.GetNodeName(opType)); + node.AddAttribute("type", LoaderSignature); + node.AddAttribute("to", (byte)ex.Kind); + if (ex.HasKeyRange) + { + var key = ex.TypeDst.ItemType.AsKey; + node.AddAttribute("min", key.Min); + node.AddAttribute("max", key.Count); + node.AddAttribute("contiguous", key.Contiguous); + } + } + + return true; + } + private static bool TryCreateEx(IExceptionContext ectx, ColInfo info, DataKind kind, KeyRange range, out PrimitiveType itemType, out ColInfoEx ex) { ectx.AssertValue(info); diff --git a/src/Microsoft.ML.Data/Transforms/KeyToVectorTransform.cs b/src/Microsoft.ML.Data/Transforms/KeyToVectorTransform.cs index 70afe195b0..05e290649b 100644 --- a/src/Microsoft.ML.Data/Transforms/KeyToVectorTransform.cs +++ b/src/Microsoft.ML.Data/Transforms/KeyToVectorTransform.cs @@ -608,7 +608,7 @@ private ValueGetter> MakeGetterInd(IRow input, int iinfo) }; } - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public bool CanSavePfa => true; diff --git a/src/Microsoft.ML.Data/Transforms/NormalizeColumn.cs b/src/Microsoft.ML.Data/Transforms/NormalizeColumn.cs index 69021d30a3..c7d7e16d33 100644 --- a/src/Microsoft.ML.Data/Transforms/NormalizeColumn.cs +++ b/src/Microsoft.ML.Data/Transforms/NormalizeColumn.cs @@ -384,7 +384,7 @@ private AffineColumnFunction(IHost host) public abstract void Save(ModelSaveContext ctx); public abstract JToken PfaInfo(BoundPfaContext ctx, JToken srcToken); - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public abstract bool OnnxInfo(OnnxContext ctx, OnnxNode nodeProtoWrapper, int featureCount); public abstract Delegate GetGetter(IRow input, int icol); @@ -503,7 +503,7 @@ public JToken PfaInfo(BoundPfaContext ctx, JToken srcToken) return null; } - public bool CanSaveOnnx => false; + public bool CanSaveOnnx(OnnxContext ctx) => false; public bool OnnxInfo(OnnxContext ctx, OnnxNode nodeProtoWrapper, int featureCount) => throw Host.ExceptNotSupp(); @@ -636,7 +636,7 @@ public JToken PfaInfo(BoundPfaContext ctx, JToken srcToken) return null; } - public bool CanSaveOnnx => false; + public bool CanSaveOnnx(OnnxContext ctx) => false; public bool OnnxInfo(OnnxContext ctx, OnnxNode nodeProtoWrapper, int featureCount) => throw Host.ExceptNotSupp(); diff --git a/src/Microsoft.ML.Data/Transforms/NormalizeUtils.cs b/src/Microsoft.ML.Data/Transforms/NormalizeUtils.cs index ab79ab4e11..26e80d150a 100644 --- a/src/Microsoft.ML.Data/Transforms/NormalizeUtils.cs +++ b/src/Microsoft.ML.Data/Transforms/NormalizeUtils.cs @@ -60,7 +60,7 @@ internal interface IColumnFunction : ICanSaveModel JToken PfaInfo(BoundPfaContext ctx, JToken srcToken); - bool CanSaveOnnx { get; } + bool CanSaveOnnx(OnnxContext ctx); bool OnnxInfo(OnnxContext ctx, OnnxNode nodeProtoWrapper, int featureCount); } diff --git a/src/Microsoft.ML.Data/Transforms/Normalizer.cs b/src/Microsoft.ML.Data/Transforms/Normalizer.cs index 8f7c25166b..6158444b24 100644 --- a/src/Microsoft.ML.Data/Transforms/Normalizer.cs +++ b/src/Microsoft.ML.Data/Transforms/Normalizer.cs @@ -452,7 +452,7 @@ private sealed class Mapper : MapperBase, ISaveAsOnnx, ISaveAsPfa { private NormalizerTransformer _parent; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public bool CanSavePfa => true; public Mapper(NormalizerTransformer parent, ISchema schema) @@ -562,12 +562,12 @@ private bool SaveAsOnnxCore(OnnxContext ctx, int iinfo, ColumnInfo info, string Contracts.AssertValue(ctx); Contracts.Assert(0 <= iinfo && iinfo < _parent._columns.Length); Contracts.Assert(_parent._columns[iinfo] == info); - Contracts.Assert(CanSaveOnnx); + Contracts.Assert(CanSaveOnnx(ctx)); if (info.InputType.ValueCount == 0) return false; - if (info.ColumnFunction.CanSaveOnnx) + if (info.ColumnFunction.CanSaveOnnx(ctx)) { string opType = "Scaler"; var node = ctx.CreateNode(opType, srcVariableName, dstVariableName, ctx.GetNodeName(opType)); diff --git a/src/Microsoft.ML.Data/Transforms/TermTransform.cs b/src/Microsoft.ML.Data/Transforms/TermTransform.cs index 5af1d970c8..e2957b8225 100644 --- a/src/Microsoft.ML.Data/Transforms/TermTransform.cs +++ b/src/Microsoft.ML.Data/Transforms/TermTransform.cs @@ -727,7 +727,7 @@ private sealed class Mapper : MapperBase, ISaveAsOnnx, ISaveAsPfa private readonly BoundTermMap[] _termMap; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; public bool CanSavePfa => true; diff --git a/src/Microsoft.ML.Data/Transforms/TransformBase.cs b/src/Microsoft.ML.Data/Transforms/TransformBase.cs index 2d9cedb17b..53f866424b 100644 --- a/src/Microsoft.ML.Data/Transforms/TransformBase.cs +++ b/src/Microsoft.ML.Data/Transforms/TransformBase.cs @@ -469,7 +469,7 @@ private sealed class ColumnTmp : OneToOneColumn public virtual bool CanSavePfa => false; - public virtual bool CanSaveOnnx => false; + public virtual bool CanSaveOnnx(OnnxContext ctx) => false; protected OneToOneTransformBase(IHostEnvironment env, string name, OneToOneColumn[] column, IDataView input, Func testType) @@ -574,7 +574,7 @@ public void SaveAsPfa(BoundPfaContext ctx) public void SaveAsOnnx(OnnxContext ctx) { Host.CheckValue(ctx, nameof(ctx)); - Host.Assert(CanSaveOnnx); + Host.Assert(CanSaveOnnx(ctx)); for (int iinfo = 0; iinfo < Infos.Length; ++iinfo) { diff --git a/src/Microsoft.ML.FastTree/FastTree.cs b/src/Microsoft.ML.FastTree/FastTree.cs index 8e5d48f260..a57d8bcc49 100644 --- a/src/Microsoft.ML.FastTree/FastTree.cs +++ b/src/Microsoft.ML.FastTree/FastTree.cs @@ -2806,7 +2806,7 @@ public abstract class FastTreePredictionWrapper : public ColumnType InputType { get; } public ColumnType OutputType => NumberType.Float; public bool CanSavePfa => true; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; protected FastTreePredictionWrapper(IHostEnvironment env, string name, Ensemble trainedEnsemble, int numFeatures, string innerArgs) : base(env, name) diff --git a/src/Microsoft.ML.Onnx/OnnxContextImpl.cs b/src/Microsoft.ML.Onnx/OnnxContextImpl.cs index 5341b35d55..7882cf3922 100644 --- a/src/Microsoft.ML.Onnx/OnnxContextImpl.cs +++ b/src/Microsoft.ML.Onnx/OnnxContextImpl.cs @@ -31,9 +31,10 @@ internal sealed class OnnxContextImpl : OnnxContext private readonly string _domain; private readonly string _producerVersion; private readonly long _modelVersion; + private readonly OnnxVersion _onnxVersion; public OnnxContextImpl(IHostEnvironment env, string name, string producerName, - string producerVersion, long modelVersion, string domain) + string producerVersion, long modelVersion, string domain, OnnxVersion onnxVersion) { Contracts.CheckValue(env, nameof(env)); _host = env.Register(nameof(OnnxContext)); @@ -52,6 +53,7 @@ public OnnxContextImpl(IHostEnvironment env, string name, string producerName, _producerVersion = producerVersion; _modelVersion = modelVersion; _domain = domain; + _onnxVersion = onnxVersion; } public override bool ContainsColumn(string colName) => _columnNameMap.ContainsKey(colName); @@ -251,5 +253,7 @@ public void AddInputVariable(ColumnType type, string colName) /// public ModelProto MakeModel() => OnnxUtils.MakeModel(_nodes, _producerName, _name, _domain, _producerVersion, _modelVersion, _inputs, _outputs, _intermediateValues); + + public override OnnxVersion GetOnnxVersion() => _onnxVersion; } } diff --git a/src/Microsoft.ML.Onnx/OnnxUtils.cs b/src/Microsoft.ML.Onnx/OnnxUtils.cs index 9605226846..8a18e56987 100644 --- a/src/Microsoft.ML.Onnx/OnnxUtils.cs +++ b/src/Microsoft.ML.Onnx/OnnxUtils.cs @@ -307,14 +307,38 @@ public static ModelArgs GetModelArgs(ColumnType type, string colName, case DataKind.TX: dataType = TensorProto.Types.DataType.String; break; + case DataKind.I1: + dataType = TensorProto.Types.DataType.Int8; + break; + case DataKind.U1: + dataType = TensorProto.Types.DataType.Uint8; + break; + case DataKind.I2: + dataType = TensorProto.Types.DataType.Int16; + break; + case DataKind.U2: + dataType = TensorProto.Types.DataType.Uint16; + break; + case DataKind.I4: + dataType = TensorProto.Types.DataType.Int32; + break; case DataKind.U4: dataType = TensorProto.Types.DataType.Int64; break; + case DataKind.I8: + dataType = TensorProto.Types.DataType.Int64; + break; + case DataKind.U8: + dataType = TensorProto.Types.DataType.Uint64; + break; case DataKind.R4: dataType = TensorProto.Types.DataType.Float; break; + case DataKind.R8: + dataType = TensorProto.Types.DataType.Double; + break; default: - Contracts.Assert(false, "Unknown type."); + Contracts.Assert(false, "Unsupported type: DataKind " + rawKind.ToString()); break; } diff --git a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs index b68f22b919..23713573cc 100644 --- a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs +++ b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs @@ -57,6 +57,9 @@ public sealed class Arguments : DataCommand.ArgumentsBase [Argument(ArgumentType.Required, Visibility = ArgumentAttribute.VisibilityType.EntryPointsOnly, HelpText = "Model that needs to be converted to ONNX format.", SortOrder = 10)] public ITransformModel Model; + + [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"latest\" or \"experimental\"", SortOrder = 11)] + public OnnxVersion OnnxVersion; } private readonly string _outputModelPath; @@ -111,7 +114,7 @@ public override void Run() } } - private void GetPipe(IChannel ch, IDataView end, out IDataView source, out IDataView trueEnd, out LinkedList transforms) + private void GetPipe(OnnxContextImpl ctx, IChannel ch, IDataView end, out IDataView source, out IDataView trueEnd, out LinkedList transforms) { Host.AssertValue(end); source = trueEnd = (end as CompositeDataLoader)?.View ?? end; @@ -120,7 +123,7 @@ private void GetPipe(IChannel ch, IDataView end, out IDataView source, out IData while (transform != null) { ITransformCanSaveOnnx onnxTransform = transform as ITransformCanSaveOnnx; - if (onnxTransform == null || !onnxTransform.CanSaveOnnx) + if (onnxTransform == null || !onnxTransform.CanSaveOnnx(ctx)) { ch.Warning("Had to stop walkback of pipeline at {0} since it cannot save itself as ONNX.", transform.GetType().Name); while (source as IDataTransform != null) @@ -160,18 +163,19 @@ private void Run(IChannel ch) else view = _model.Apply(Host, new EmptyDataView(Host, _model.InputSchema)); + // Create the ONNX context for storing global information + var assembly = System.Reflection.Assembly.GetExecutingAssembly(); + var versionInfo = System.Diagnostics.FileVersionInfo.GetVersionInfo(assembly.Location); + var ctx = new OnnxContextImpl(Host, _name, ProducerName, versionInfo.FileVersion, + ModelVersion, _domain, Args.OnnxVersion); + // Get the transform chain. IDataView source; IDataView end; LinkedList transforms; - GetPipe(ch, view, out source, out end, out transforms); + GetPipe(ctx, ch, view, out source, out end, out transforms); Host.Assert(transforms.Count == 0 || transforms.Last.Value == end); - var assembly = System.Reflection.Assembly.GetExecutingAssembly(); - var versionInfo = System.Diagnostics.FileVersionInfo.GetVersionInfo(assembly.Location); - - var ctx = new OnnxContextImpl(Host, _name, ProducerName, versionInfo.FileVersion, - ModelVersion, _domain); // If we have a predictor, try to get the scorer for it. if (rawPred != null) { @@ -188,7 +192,7 @@ private void Run(IChannel ch) var scorePipe = ScoreUtils.GetScorer(rawPred, data, Host, trainSchema); var scoreOnnx = scorePipe as ITransformCanSaveOnnx; - if (scoreOnnx?.CanSaveOnnx == true) + if (scoreOnnx?.CanSaveOnnx(ctx) == true) { Host.Assert(scorePipe.Source == end); end = scorePipe; @@ -222,7 +226,7 @@ private void Run(IChannel ch) //Create graph nodes, outputs and intermediate values. foreach (var trans in transforms) { - Host.Assert(trans.CanSaveOnnx); + Host.Assert(trans.CanSaveOnnx(ctx)); trans.SaveAsOnnx(ctx); } diff --git a/src/Microsoft.ML.StandardLearners/Standard/LinearPredictor.cs b/src/Microsoft.ML.StandardLearners/Standard/LinearPredictor.cs index 2a5d73705f..284add3789 100644 --- a/src/Microsoft.ML.StandardLearners/Standard/LinearPredictor.cs +++ b/src/Microsoft.ML.StandardLearners/Standard/LinearPredictor.cs @@ -101,7 +101,7 @@ IEnumerator IEnumerable.GetEnumerator() public bool CanSavePfa => true; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; /// /// Constructs a new linear predictor. diff --git a/src/Microsoft.ML.StandardLearners/Standard/LogisticRegression/MulticlassLogisticRegression.cs b/src/Microsoft.ML.StandardLearners/Standard/LogisticRegression/MulticlassLogisticRegression.cs index ca850fe46e..524a2656e0 100644 --- a/src/Microsoft.ML.StandardLearners/Standard/LogisticRegression/MulticlassLogisticRegression.cs +++ b/src/Microsoft.ML.StandardLearners/Standard/LogisticRegression/MulticlassLogisticRegression.cs @@ -331,7 +331,7 @@ private static VersionInfo GetVersionInfo() public ColumnType InputType { get; } public ColumnType OutputType { get; } public bool CanSavePfa => true; - public bool CanSaveOnnx => true; + public bool CanSaveOnnx(OnnxContext ctx) => true; internal MulticlassLogisticRegressionPredictor(IHostEnvironment env, ref VBuffer weights, int numClasses, int numFeatures, string[] labelNames, LinearModelStatistics stats = null) : base(env, RegistrationName) diff --git a/src/Microsoft.ML.Transforms/NAReplaceTransform.cs b/src/Microsoft.ML.Transforms/NAReplaceTransform.cs index c9b309af89..6c50aa805d 100644 --- a/src/Microsoft.ML.Transforms/NAReplaceTransform.cs +++ b/src/Microsoft.ML.Transforms/NAReplaceTransform.cs @@ -183,7 +183,7 @@ private static string TestType(ColumnType type) // The isNA delegates, parallel to Infos. private readonly Delegate[] _isNAs; - public override bool CanSaveOnnx => true; + public override bool CanSaveOnnx(OnnxContext ctx) => true; /// /// Convenience constructor for public facing API. From 36aba972f4549ddd8a4d7d759cdddac1a3ea4cff Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Wed, 19 Sep 2018 11:45:53 -0700 Subject: [PATCH 2/7] Update CSharpAPI and entry point --- src/Microsoft.ML.Legacy/CSharpApi.cs | 11 +++++++++++ .../Common/EntryPoints/core_manifest.json | 15 +++++++++++++++ 2 files changed, 26 insertions(+) diff --git a/src/Microsoft.ML.Legacy/CSharpApi.cs b/src/Microsoft.ML.Legacy/CSharpApi.cs index ba6f0f866e..e9dbc6e622 100644 --- a/src/Microsoft.ML.Legacy/CSharpApi.cs +++ b/src/Microsoft.ML.Legacy/CSharpApi.cs @@ -3266,6 +3266,12 @@ public OneVersusAllPipelineStep(Output output) namespace Legacy.Models { + public enum OnnxVersion + { + Latest = 0, + Experimental = 1 + } + /// /// Converts the model to ONNX format. @@ -3309,6 +3315,11 @@ public sealed partial class OnnxConverter /// public Var Model { get; set; } = new Var(); + /// + /// The targeted ONNX version. It can be either "latest" or "experimental" + /// + public OnnxVersion OnnxVersion { get; set; } = OnnxVersion.Latest; + /// /// The data file /// diff --git a/test/BaselineOutput/Common/EntryPoints/core_manifest.json b/test/BaselineOutput/Common/EntryPoints/core_manifest.json index 1ff41a090b..ab28b0e98e 100644 --- a/test/BaselineOutput/Common/EntryPoints/core_manifest.json +++ b/test/BaselineOutput/Common/EntryPoints/core_manifest.json @@ -2479,6 +2479,21 @@ "Required": true, "SortOrder": 10.0, "IsNullable": false + }, + { + "Name": "OnnxVersion", + "Type": { + "Kind": "Enum", + "Values": [ + "Latest", + "Experimental" + ] + }, + "Desc": "The targeted ONNX version. It can be either \"latest\" or \"experimental\"", + "Required": false, + "SortOrder": 11.0, + "IsNullable": false, + "Default": "Latest" } ], "Outputs": [] From 57a942948cf1010e711601b56bc7c01a35968ee5 Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Mon, 24 Sep 2018 17:37:21 -0700 Subject: [PATCH 3/7] Address one comment --- src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs | 4 ++-- src/Microsoft.ML.Onnx/SaveOnnxCommand.cs | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs b/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs index 9dc62ac55b..423ede8b8f 100644 --- a/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs +++ b/src/Microsoft.ML.Data/Model/Onnx/OnnxContext.cs @@ -7,7 +7,7 @@ namespace Microsoft.ML.Runtime.Model.Onnx { - public enum OnnxVersion { Latest, Experimental } + public enum OnnxVersion { Stable=0, Experimental=1 } /// /// A context for defining a ONNX output. The context internally contains the model-in-progress being built. This @@ -102,7 +102,7 @@ public OnnxNode CreateNode(string opType, string input, string output, string na => CreateNode(opType, new[] { input }, new[] { output }, name, domain); /// - /// Get the targeted ONNX version string. Only two values are allowed now: "latest" and "experimental". + /// Get the targeted ONNX version string. Only two values are allowed now: "Stable" and "Experimental". /// /// public abstract OnnxVersion GetOnnxVersion(); diff --git a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs index 23713573cc..44d3ee0723 100644 --- a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs +++ b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs @@ -58,7 +58,7 @@ public sealed class Arguments : DataCommand.ArgumentsBase [Argument(ArgumentType.Required, Visibility = ArgumentAttribute.VisibilityType.EntryPointsOnly, HelpText = "Model that needs to be converted to ONNX format.", SortOrder = 10)] public ITransformModel Model; - [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"latest\" or \"experimental\"", SortOrder = 11)] + [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\"", SortOrder = 11)] public OnnxVersion OnnxVersion; } From 3df5fca5bae9a0689cdb5684b93d24d8a0edd768 Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Tue, 25 Sep 2018 14:10:26 -0700 Subject: [PATCH 4/7] Update old APIs to reflect enum's change --- src/Microsoft.ML.Legacy/CSharpApi.cs | 6 +++--- test/BaselineOutput/Common/EntryPoints/core_manifest.json | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/src/Microsoft.ML.Legacy/CSharpApi.cs b/src/Microsoft.ML.Legacy/CSharpApi.cs index 38846217da..c960a56314 100644 --- a/src/Microsoft.ML.Legacy/CSharpApi.cs +++ b/src/Microsoft.ML.Legacy/CSharpApi.cs @@ -3268,7 +3268,7 @@ namespace Legacy.Models { public enum OnnxVersion { - Latest = 0, + Stable = 0, Experimental = 1 } @@ -3316,9 +3316,9 @@ public sealed partial class OnnxConverter public Var Model { get; set; } = new Var(); /// - /// The targeted ONNX version. It can be either "latest" or "experimental" + /// The targeted ONNX version. It can be either "Stable" or "Experimental" /// - public OnnxVersion OnnxVersion { get; set; } = OnnxVersion.Latest; + public OnnxVersion OnnxVersion { get; set; } = OnnxVersion.Stable; /// /// The data file diff --git a/test/BaselineOutput/Common/EntryPoints/core_manifest.json b/test/BaselineOutput/Common/EntryPoints/core_manifest.json index 508b67349b..f01187e3c6 100644 --- a/test/BaselineOutput/Common/EntryPoints/core_manifest.json +++ b/test/BaselineOutput/Common/EntryPoints/core_manifest.json @@ -2485,15 +2485,15 @@ "Type": { "Kind": "Enum", "Values": [ - "Latest", + "Stable", "Experimental" ] }, - "Desc": "The targeted ONNX version. It can be either \"latest\" or \"experimental\"", + "Desc": "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\"", "Required": false, "SortOrder": 11.0, "IsNullable": false, - "Default": "Latest" + "Default": "Stable" } ], "Outputs": [] From ac2b29d56144a9c87bfdcfc9880c87bb8a1334af Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Wed, 26 Sep 2018 16:47:34 -0700 Subject: [PATCH 5/7] Extend doc string for targeted version --- src/Microsoft.ML.Legacy/CSharpApi.cs | 2 +- src/Microsoft.ML.Onnx/SaveOnnxCommand.cs | 2 +- test/BaselineOutput/Common/EntryPoints/core_manifest.json | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/Microsoft.ML.Legacy/CSharpApi.cs b/src/Microsoft.ML.Legacy/CSharpApi.cs index c960a56314..e35101ccb4 100644 --- a/src/Microsoft.ML.Legacy/CSharpApi.cs +++ b/src/Microsoft.ML.Legacy/CSharpApi.cs @@ -3316,7 +3316,7 @@ public sealed partial class OnnxConverter public Var Model { get; set; } = new Var(); /// - /// The targeted ONNX version. It can be either "Stable" or "Experimental" + /// The targeted ONNX version. It can be either "Stable" or "Experimental". If "Experimentab" is used, models produced can contain components not officially defined in ONNX standard. /// public OnnxVersion OnnxVersion { get; set; } = OnnxVersion.Stable; diff --git a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs index 44d3ee0723..b6488cb55c 100644 --- a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs +++ b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs @@ -58,7 +58,7 @@ public sealed class Arguments : DataCommand.ArgumentsBase [Argument(ArgumentType.Required, Visibility = ArgumentAttribute.VisibilityType.EntryPointsOnly, HelpText = "Model that needs to be converted to ONNX format.", SortOrder = 10)] public ITransformModel Model; - [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\"", SortOrder = 11)] + [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimentab\" is used, models produced can contain components not officially defined in ONNX standard.", SortOrder = 11)] public OnnxVersion OnnxVersion; } diff --git a/test/BaselineOutput/Common/EntryPoints/core_manifest.json b/test/BaselineOutput/Common/EntryPoints/core_manifest.json index f01187e3c6..e7c0174566 100644 --- a/test/BaselineOutput/Common/EntryPoints/core_manifest.json +++ b/test/BaselineOutput/Common/EntryPoints/core_manifest.json @@ -2489,7 +2489,7 @@ "Experimental" ] }, - "Desc": "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\"", + "Desc": "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimentab\" is used, models produced can contain components not officially defined in ONNX standard.", "Required": false, "SortOrder": 11.0, "IsNullable": false, From 845f7e960badc3c618ad8e6598fb6dcfe66ac1f0 Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Tue, 9 Oct 2018 10:36:57 -0700 Subject: [PATCH 6/7] Address comments --- src/Microsoft.ML.Onnx/OnnxUtils.cs | 3 ++- src/Microsoft.ML.Onnx/SaveOnnxCommand.cs | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/src/Microsoft.ML.Onnx/OnnxUtils.cs b/src/Microsoft.ML.Onnx/OnnxUtils.cs index d36c700c51..83623ab0ba 100644 --- a/src/Microsoft.ML.Onnx/OnnxUtils.cs +++ b/src/Microsoft.ML.Onnx/OnnxUtils.cs @@ -342,7 +342,8 @@ public static ModelArgs GetModelArgs(ColumnType type, string colName, dataType = TensorProto.Types.DataType.Double; break; default: - Contracts.Assert(false, "Unsupported type: DataKind " + rawKind.ToString()); + string msg = "Unsupported type: DataKind " + rawKind.ToString(); + Contracts.Check(false, msg); break; } diff --git a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs index 5eda42781d..fca84fb8b3 100644 --- a/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs +++ b/src/Microsoft.ML.Onnx/SaveOnnxCommand.cs @@ -58,7 +58,7 @@ public sealed class Arguments : DataCommand.ArgumentsBase [Argument(ArgumentType.Required, Visibility = ArgumentAttribute.VisibilityType.EntryPointsOnly, HelpText = "Model that needs to be converted to ONNX format.", SortOrder = 10)] public ITransformModel Model; - [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimental\" is used, models produced can contain components not officially defined in ONNX standard.", SortOrder = 11)] + [Argument(ArgumentType.AtMostOnce, HelpText = "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimental\" is used, produced model can contain components that is not officially supported in ONNX standard.", SortOrder = 11)] public OnnxVersion OnnxVersion; } From 60e3f332d3f77f79676ff02356f3118897f4db00 Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Tue, 9 Oct 2018 15:49:19 -0700 Subject: [PATCH 7/7] Update API --- src/Microsoft.ML.Legacy/CSharpApi.cs | 2 +- test/BaselineOutput/Common/EntryPoints/core_manifest.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/Microsoft.ML.Legacy/CSharpApi.cs b/src/Microsoft.ML.Legacy/CSharpApi.cs index 7bc340cef6..42518890a5 100644 --- a/src/Microsoft.ML.Legacy/CSharpApi.cs +++ b/src/Microsoft.ML.Legacy/CSharpApi.cs @@ -3412,7 +3412,7 @@ public sealed partial class OnnxConverter public Var Model { get; set; } = new Var(); /// - /// The targeted ONNX version. It can be either "Stable" or "Experimental". If "Experimental" is used, models produced can contain components not officially defined in ONNX standard. + /// The targeted ONNX version. It can be either "Stable" or "Experimental". If "Experimental" is used, produced model can contain components that is not officially supported in ONNX standard. /// public OnnxVersion OnnxVersion { get; set; } = OnnxVersion.Stable; diff --git a/test/BaselineOutput/Common/EntryPoints/core_manifest.json b/test/BaselineOutput/Common/EntryPoints/core_manifest.json index 6f83c82a4d..edd0fa93f4 100644 --- a/test/BaselineOutput/Common/EntryPoints/core_manifest.json +++ b/test/BaselineOutput/Common/EntryPoints/core_manifest.json @@ -2489,7 +2489,7 @@ "Experimental" ] }, - "Desc": "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimental\" is used, models produced can contain components not officially defined in ONNX standard.", + "Desc": "The targeted ONNX version. It can be either \"Stable\" or \"Experimental\". If \"Experimental\" is used, produced model can contain components that is not officially supported in ONNX standard.", "Required": false, "SortOrder": 11.0, "IsNullable": false,