From a19fa4a3f0c1b6d7f5c98bd80dd9b98e0c40b666 Mon Sep 17 00:00:00 2001 From: "Harish S. Kulkarni" Date: Tue, 25 Feb 2020 12:57:28 -0800 Subject: [PATCH] Fixed bugs in OptionalColumnTransform and ColumnSelecting --- src/Microsoft.ML.Data/Transforms/ColumnSelecting.cs | 7 +++---- src/Microsoft.ML.Transforms/OptionalColumnTransform.cs | 7 ++++++- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/src/Microsoft.ML.Data/Transforms/ColumnSelecting.cs b/src/Microsoft.ML.Data/Transforms/ColumnSelecting.cs index 669b89cc27..65a736b1e5 100644 --- a/src/Microsoft.ML.Data/Transforms/ColumnSelecting.cs +++ b/src/Microsoft.ML.Data/Transforms/ColumnSelecting.cs @@ -732,19 +732,18 @@ IDataTransform ITransformTemplate.ApplyToData(IHostEnvironment env, IDataView ne public void SaveAsOnnx(OnnxContext ctx) { - var droppedCols = new HashSet(Enumerable.Range(0, InputSchema.Count)); - var outputToInputMap = _mapper.OutputToInputMap; for(int i = 0; i < outputToInputMap.Length; i++) { var srcCol = InputSchema[outputToInputMap[i]]; var dstCol = OutputSchema[i]; + if (!ctx.ContainsColumn(srcCol.Name) || dstCol.IsHidden) + continue; + var srcVariable = ctx.GetVariableName(srcCol.Name); var dstVariable = ctx.AddIntermediateVariable(dstCol.Type, dstCol.Name, true); string opType = "Identity"; ctx.CreateNode(opType, srcVariable, dstVariable, ctx.GetNodeName(opType), ""); - - droppedCols.Remove(srcCol.Index); } } } diff --git a/src/Microsoft.ML.Transforms/OptionalColumnTransform.cs b/src/Microsoft.ML.Transforms/OptionalColumnTransform.cs index 201abc1222..97cb6d4a3f 100644 --- a/src/Microsoft.ML.Transforms/OptionalColumnTransform.cs +++ b/src/Microsoft.ML.Transforms/OptionalColumnTransform.cs @@ -511,7 +511,12 @@ public void SaveAsOnnx(OnnxContext ctx) if (!ctx.ContainsColumn(inputColumnName)) continue; - if (!SaveAsOnnxCore(ctx, ctx.GetVariableName(inputColumnName), _bindings.ColumnTypes[iinfo])) + // If there is already a column of this name, don't add this column as an OptionalColumn/Initializer + var srcVariableName = ctx.GetVariableName(inputColumnName); + if (srcVariableName != inputColumnName) + continue; + + if (!SaveAsOnnxCore(ctx, srcVariableName, _bindings.ColumnTypes[iinfo])) ctx.RemoveColumn(inputColumnName, true); } }