diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index d24965578d..4dd77fd822 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -495,13 +495,14 @@ internal static IEnumerable CrossValidationSplit(IHostEnvironment /// internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) { + Contracts.CheckValue(env, nameof(env)); + var host = env.Register("rand"); // We need to handle two cases: if samplingKeyColumn is provided, we use hashJoin to // build a single hash of it. If it is not, we generate a random number. - if (samplingKeyColumn == null) { samplingKeyColumn = data.Schema.GetTempColumnName("SamplingKeyColumn"); - data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)seed); + data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)(seed ?? host.Rand.Next())); } else { @@ -517,11 +518,7 @@ internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDa // instead of having two hash transformations. var origStratCol = samplingKeyColumn; samplingKeyColumn = data.Schema.GetTempColumnName(samplingKeyColumn); - HashingEstimator.ColumnOptionsInternal columnOptions; - if (seed.HasValue) - columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)seed.Value); - else - columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30); + var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? host.Rand.Next())); data = new HashingEstimator(env, columnOptions).Fit(data).Transform(data); } else @@ -533,7 +530,6 @@ internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDa data = new NormalizingEstimator(env, new NormalizingEstimator.MinMaxColumnOptions(samplingKeyColumn, origStratCol, ensureZeroUntouched: true)).Fit(data).Transform(data); } } - } } }