Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 41 additions & 48 deletions src/Microsoft.ML.Data/EntryPoints/EntryPointNode.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -475,7 +475,7 @@ public float Cost

private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCatalog, RunContext context,
string id, string entryPointName, JObject inputs, JObject outputs, bool checkpoint = false,
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null)
string stageId = "", float cost = float.NaN, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertNonEmpty(id);
Expand DownExpand Up@@ -510,49 +510,10 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
throw _host.Except($"The following required inputs were not provided: {String.Join(", ", missing)}");

var inputInstance = _inputBuilder.GetInstance();
var warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(label) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithLabel)))
{
var labelColField = _inputBuilder.GetFieldNameOrNull("LabelColumn");
ch.AssertNonEmpty(labelColField);
var labelColFieldType = _inputBuilder.GetFieldTypeOrNull(labelColField);
ch.Assert(labelColFieldType == typeof(string));
var inputLabel = inputInstance.GetType().GetField(labelColField).GetValue(inputInstance);
if (label != (string)inputLabel)
ch.Warning(warning, "label", label, inputLabel);
else
_inputBuilder.TrySetValue(labelColField, label);
}
if (!string.IsNullOrEmpty(group) && Utils.Size(_entryPoint.InputKinds) > 0 &&
_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithGroupId)))
{
var groupColField = _inputBuilder.GetFieldNameOrNull("GroupIdColumn");
ch.AssertNonEmpty(groupColField);
var groupColFieldType = _inputBuilder.GetFieldTypeOrNull(groupColField);
ch.Assert(groupColFieldType == typeof(string));
var inputGroup = inputInstance.GetType().GetField(groupColField).GetValue(inputInstance);
if (group != (Optional<string>)inputGroup)
ch.Warning(warning, "group Id", label, inputGroup);
else
_inputBuilder.TrySetValue(groupColField, label);
}
if (!string.IsNullOrEmpty(weight) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(_entryPoint.InputKinds.Contains(typeof(CommonInputs.ITrainerInputWithWeight)) ||
_entryPoint.InputKinds.Contains(typeof(CommonInputs.IUnsupervisedTrainerWithWeight))))
{
var weightColField = _inputBuilder.GetFieldNameOrNull("WeightColumn");
ch.AssertNonEmpty(weightColField);
var weightColFieldType = _inputBuilder.GetFieldTypeOrNull(weightColField);
ch.Assert(weightColFieldType == typeof(string));
var inputWeight = inputInstance.GetType().GetField(weightColField).GetValue(inputInstance);
if (weight != (Optional<string>)inputWeight)
ch.Warning(warning, "weight", label, inputWeight);
else
_inputBuilder.TrySetValue(weightColField, label);
}
SetColumnArgument(ch, inputInstance, "LabelColumn", label, "label", typeof(CommonInputs.ITrainerInputWithLabel));
SetColumnArgument(ch, inputInstance, "GroupIdColumn", group, "group Id", typeof(CommonInputs.ITrainerInputWithGroupId));
SetColumnArgument(ch, inputInstance, "WeightColumn", weight, "weight", typeof(CommonInputs.ITrainerInputWithWeight), typeof(CommonInputs.IUnsupervisedTrainerWithWeight));
SetColumnArgument(ch, inputInstance, "NameColumn", name, "name");

// Validate outputs.
_outputHelper = new OutputHelper(_host, _entryPoint.OutputType);
Expand All@@ -568,6 +529,38 @@ private EntryPointNode(IHostEnvironment env, IChannel ch, ModuleCatalog moduleCa
Cost = cost;
}

private void SetColumnArgument(IChannel ch, object inputInstance, string argName, string colName, string columnRole, params Type[] inputKinds)
{
Contracts.AssertValue(ch);
ch.AssertValue(inputInstance);
ch.AssertNonEmpty(argName);
ch.AssertValueOrNull(colName);
ch.AssertNonEmpty(columnRole);
ch.AssertValueOrNull(inputKinds);

var colField = _inputBuilder.GetFieldNameOrNull(argName);
if (string.IsNullOrEmpty(colField))
return;

const string warning = "Different {0} column specified in trainer and in macro: '{1}', '{2}'." +
" Using column '{2}'. To column use '{1}' instead, please specify this name in" +
"the trainer node arguments.";
if (!string.IsNullOrEmpty(colName) && Utils.Size(_entryPoint.InputKinds) > 0 &&
(Utils.Size(inputKinds) == 0 || _entryPoint.InputKinds.Intersect(inputKinds).Any()))
{
ch.AssertNonEmpty(colField);
var colFieldType = _inputBuilder.GetFieldTypeOrNull(colField);
ch.Assert(colFieldType == typeof(string));
var inputColName = inputInstance.GetType().GetField(colField).GetValue(inputInstance);
ch.Assert(inputColName is string || inputColName is Optional<string>);
var str = inputColName is string ? (string)inputColName : ((Optional<string>)inputColName).Value;
if (colName != str)
ch.Warning(warning, columnRole, colName, inputColName);
else
_inputBuilder.TrySetValue(colField, colName);
}
}

public static EntryPointNode Create(
IHostEnvironment env,
string entryPointName,
Expand DownExpand Up@@ -902,7 +895,7 @@ private object BuildParameterValue(List<ParameterBinding> bindings)
}

public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContext context, JArray nodes,
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null)
ModuleCatalog moduleCatalog, string label = null, string group = null, string weight = null, string name = null)
{
Contracts.AssertValue(env);
env.AssertValue(context);
Expand All@@ -918,7 +911,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (node == null)
throw env.Except("Unexpected node token: '{0}'", nodes[i]);

string name = node[FieldNames.Name].Value<string>();
string nodeName = node[FieldNames.Name].Value<string>();
var inputs = node[FieldNames.Inputs] as JObject;
if (inputs == null && node[FieldNames.Inputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Inputs, node[FieldNames.Inputs]);
Expand All@@ -927,7 +920,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
if (outputs == null && node[FieldNames.Outputs] != null)
throw env.Except("Unexpected {0} token: '{1}'", FieldNames.Outputs, node[FieldNames.Outputs]);

var id = context.GenerateId(name);
var id = context.GenerateId(nodeName);
var unexpectedFields = node.Properties().Where(
x => x.Name != FieldNames.Name && x.Name != FieldNames.Inputs && x.Name != FieldNames.Outputs
&& x.Name != FieldNames.StageId && x.Name != FieldNames.Checkpoint && x.Name != FieldNames.Cost);
Expand All@@ -942,7 +935,7 @@ public static List<EntryPointNode> ValidateNodes(IHostEnvironment env, RunContex
ch.Warning("Node '{0}' has unexpected fields that are ignored: {1}", id, string.Join(", ", unexpectedFields.Select(x => x.Name)));
}

result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, name, inputs, outputs, checkpoint, stageId, cost, label, group, weight));
result.Add(new EntryPointNode(env, ch, moduleCatalog, context, id, nodeName, inputs, outputs, checkpoint, stageId, cost, label, group, weight, name));
}

ch.Done();
Expand Down
15 changes: 15 additions & 0 deletions src/Microsoft.ML/CSharpApi.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -2509,6 +2509,11 @@ public sealed partial class CrossValidationResultsCombiner
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }

/// <summary>
/// Specifies the trainer kind, which determines the evaluator to be used.
/// </summary>
Expand DownExpand Up@@ -2629,6 +2634,11 @@ public sealed partial class CrossValidator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand DownExpand Up@@ -4020,6 +4030,11 @@ public sealed partial class TrainTestEvaluator
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> GroupColumn { get; set; }

/// <summary>
/// Name column name
/// </summary>
public Microsoft.ML.Runtime.EntryPoints.Optional<string> NameColumn { get; set; }


public sealed class Output
{
Expand Down
31 changes: 20 additions & 11 deletions src/Microsoft.ML/Runtime/EntryPoints/CrossValidationMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -66,26 +66,29 @@ public sealed class Arguments

// For splitting the data into folds, this column is used for grouping rows and makes sure
// that a group of rows is not split among folds.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for stratification", ShortName = "strat", SortOrder = 6)]
public string StratificationColumn;

// The number of folds to generate.
[Argument(ArgumentType.LastOccurenceWins, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Number of folds in k-fold cross-validation", ShortName = "k", SortOrder = 7)]
public int NumFolds = 2;

// REVIEW: suggest moving to subcomponents for evaluators, to allow for different parameters on the evaluators
// (and the same for the TrainTest macro). I currently do not know how to do this, so this should be revisited in the future.
[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 8)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 9)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 10)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 11)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 12)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

// REVIEW: This output would be much better as an array of CommonOutputs.ClassificationEvaluateOutput,
Expand DownExpand Up@@ -127,16 +130,19 @@ public sealed class CombineMetricsInput
[Argument(ArgumentType.Multiple, HelpText = "Warning datasets", SortOrder = 4)]
public IDataView[] Warnings;

[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 5)]
[Argument(ArgumentType.AtMostOnce, HelpText = "The label column name", ShortName = "Label", SortOrder = 6)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 7)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 8)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 6)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 9)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);

[Argument(ArgumentType.Required, HelpText = "Specifies the trainer kind, which determines the evaluator to be used.", SortOrder = 5)]
public MacroUtils.TrainerKinds Kind = MacroUtils.TrainerKinds.SignatureBinaryClassifierTrainer;
}

Expand DownExpand Up@@ -206,7 +212,8 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
TransformModel = null,
LabelColumn = input.LabelColumn,
GroupColumn = input.GroupColumn,
WeightColumn = input.WeightColumn
WeightColumn = input.WeightColumn,
NameColumn = input.NameColumn
};

if (transformModelVarName != null)
Expand DownExpand Up@@ -377,6 +384,7 @@ public static CommonOutputs.MacroOutput<Output> CrossValidate(
combineArgs.LabelColumn = input.LabelColumn;
combineArgs.WeightColumn = input.WeightColumn;
combineArgs.GroupColumn = input.GroupColumn;
combineArgs.NameColumn = input.NameColumn;

// Set the input bindings for the CombineMetrics entry point.
var combineInputBindingMap = new Dictionary<string, List<ParameterBinding>>();
Expand DownExpand Up@@ -429,7 +437,8 @@ public static CombinedOutput CombineMetrics(IHostEnvironment env, CombineMetrics
{
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Label, input.LabelColumn),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Weight, input.WeightColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value)
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Group, input.GroupColumn.Value),
RoleMappedSchema.CreatePair(RoleMappedSchema.ColumnRole.Name, input.NameColumn.Value)
})).ToArray(),
out var variableSizeVectorColumnNames);

Expand Down
16 changes: 11 additions & 5 deletions src/Microsoft.ML/Runtime/EntryPoints/TrainTestMacro.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -63,14 +63,17 @@ public sealed class Arguments
[Argument(ArgumentType.AtMostOnce, HelpText = "Indicates whether to include and output training dataset metrics.", SortOrder = 9)]
public Boolean IncludeTrainingMetrics = false;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for labels", ShortName = "lab", SortOrder = 10)]
public string LabelColumn = DefaultColumnNames.Label;

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for example weight", ShortName = "weight", SortOrder = 11)]
public Optional<string> WeightColumn = Optional<string>.Implicit(DefaultColumnNames.Weight);

[Argument(ArgumentType.LastOccurenceWins, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
[Argument(ArgumentType.AtMostOnce, HelpText = "Column to use for grouping", ShortName = "group", SortOrder = 12)]
public Optional<string> GroupColumn = Optional<string>.Implicit(DefaultColumnNames.GroupId);

[Argument(ArgumentType.AtMostOnce, HelpText = "Name column name", ShortName = "name", SortOrder = 13)]
public Optional<string> NameColumn = Optional<string>.Implicit(DefaultColumnNames.Name);
}

public sealed class Output
Expand DownExpand Up@@ -120,7 +123,9 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
// Parse the subgraph.
var subGraphRunContext = new RunContext(env);
var subGraphNodes = EntryPointNode.ValidateNodes(env, subGraphRunContext, input.Nodes, node.Catalog, input.LabelColumn,
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null, input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null);
input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
input.NameColumn.IsExplicit ? input.NameColumn.Value : null);

// Change the subgraph to use the training data as input.
var varName = input.Inputs.Data.VarName;
Expand DownExpand Up@@ -221,7 +226,8 @@ public static CommonOutputs.MacroOutput<Output> TrainTest(
{
LabelColumn = input.LabelColumn,
WeightColumn = input.WeightColumn.IsExplicit ? input.WeightColumn.Value : null,
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null
GroupColumn = input.GroupColumn.IsExplicit ? input.GroupColumn.Value : null,
NameColumn = input.NameColumn.IsExplicit ? input.NameColumn.Value : null
};

string outVariableName;
Expand Down
Loading