diff --git a/src/Microsoft.ML.Data/Data/SchemaDefinition.cs b/src/Microsoft.ML.Data/Data/SchemaDefinition.cs
index b26cd9554c..cd5cfff249 100644
--- a/src/Microsoft.ML.Data/Data/SchemaDefinition.cs
+++ b/src/Microsoft.ML.Data/Data/SchemaDefinition.cs
@@ -104,6 +104,7 @@ public sealed class ColumnNameAttribute : Attribute
///
/// Column name.
///
+ [BestFriend]
internal string Name { get; }
///
diff --git a/src/Microsoft.ML.Data/DataLoadSave/Text/LoadColumnAttribute.cs b/src/Microsoft.ML.Data/DataLoadSave/Text/LoadColumnAttribute.cs
index 60ccf5d58c..2e930cd9fc 100644
--- a/src/Microsoft.ML.Data/DataLoadSave/Text/LoadColumnAttribute.cs
+++ b/src/Microsoft.ML.Data/DataLoadSave/Text/LoadColumnAttribute.cs
@@ -46,6 +46,7 @@ public LoadColumnAttribute(int[] columnIndexes)
Sources.Add(new TextLoader.Range(col));
}
+ [BestFriend]
internal List Sources;
}
}
diff --git a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoader.cs b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoader.cs
index 2c9871fcc5..9f6cbc3c3b 100644
--- a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoader.cs
+++ b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoader.cs
@@ -5,7 +5,11 @@
using System;
using System.Collections.Generic;
using System.Data;
+using System.Data.Common;
using System.Linq;
+using System.Reflection;
+using System.Runtime.CompilerServices;
+using System.Text;
using Microsoft.ML;
using Microsoft.ML.CommandLine;
using Microsoft.ML.Data;
@@ -98,6 +102,71 @@ void ICanSaveModel.Save(ModelSaveContext ctx)
/// The source from which to load data.
public IDataView Load(DatabaseSource source) => new BoundLoader(this, source);
+ internal static DatabaseLoader CreateDatabaseLoader(IHostEnvironment host)
+ {
+ var userType = typeof(TInput);
+
+ var fieldInfos = userType.GetFields(BindingFlags.Public | BindingFlags.Instance);
+
+ var propertyInfos =
+ userType
+ .GetProperties(BindingFlags.Public | BindingFlags.Instance)
+ .Where(x => x.CanRead && x.GetGetMethod() != null && x.GetIndexParameters().Length == 0);
+
+ var memberInfos = (fieldInfos as IEnumerable).Concat(propertyInfos).ToArray();
+
+ if (memberInfos.Length == 0)
+ throw host.ExceptParam(nameof(TInput), $"Should define at least one public, readable field or property in {nameof(TInput)}.");
+
+ var columns = new List();
+
+ for (int index = 0; index < memberInfos.Length; index++)
+ {
+ var memberInfo = memberInfos[index];
+ var mappingAttrName = memberInfo.GetCustomAttribute();
+
+ var column = new Column();
+ column.Name = mappingAttrName?.Name ?? memberInfo.Name;
+
+ var mappingAttr = memberInfo.GetCustomAttribute();
+
+ if (mappingAttr is object)
+ {
+ var sources = mappingAttr.Sources.Select((source) => Range.FromTextLoaderRange(source)).ToArray();
+ column.Source = sources.Single().Min;
+ }
+
+ InternalDataKind dk;
+ switch (memberInfo)
+ {
+ case FieldInfo field:
+ if (!InternalDataKindExtensions.TryGetDataKind(field.FieldType.IsArray ? field.FieldType.GetElementType() : field.FieldType, out dk))
+ throw Contracts.Except($"Field {memberInfo.Name} is of unsupported type.");
+
+ break;
+
+ case PropertyInfo property:
+ if (!InternalDataKindExtensions.TryGetDataKind(property.PropertyType.IsArray ? property.PropertyType.GetElementType() : property.PropertyType, out dk))
+ throw Contracts.Except($"Property {memberInfo.Name} is of unsupported type.");
+ break;
+
+ default:
+ Contracts.Assert(false);
+ throw Contracts.ExceptNotSupp("Expected a FieldInfo or a PropertyInfo");
+ }
+
+ column.Type = dk.ToDbType();
+
+ columns.Add(column);
+ }
+
+ var options = new Options
+ {
+ Columns = columns.ToArray()
+ };
+ return new DatabaseLoader(host, options);
+ }
+
///
/// Describes how an input column should be mapped to an column.
///
@@ -128,6 +197,86 @@ public sealed class Column
public KeyCount KeyCount;
}
+ ///
+ /// Specifies the range of indices of input columns that should be mapped to an output column.
+ ///
+ public sealed class Range
+ {
+ public Range() { }
+
+ ///
+ /// A range representing a single value. Will result in a scalar column.
+ ///
+ /// The index of the field of the text file to read.
+ public Range(int index)
+ {
+ Contracts.CheckParam(index >= 0, nameof(index), "Must be non-negative");
+ Min = index;
+ Max = index;
+ }
+
+ ///
+ /// A range representing a set of values. Will result in a vector column.
+ ///
+ /// The minimum inclusive index of the column.
+ /// The maximum-inclusive index of the column. If null
+ /// indicates that the should auto-detect the legnth
+ /// of the lines, and read untill the end.
+ public Range(int min, int? max)
+ {
+ Contracts.CheckParam(min >= 0, nameof(min), "Must be non-negative");
+ Contracts.CheckParam(!(max < min), nameof(max), "If specified, must be greater than or equal to " + nameof(min));
+
+ Min = min;
+ Max = max;
+ // Note that without the following being set, in the case where there is a single range
+ // where Min == Max, the result will not be a vector valued but a scalar column.
+ ForceVector = true;
+ AutoEnd = max == null;
+ }
+
+ ///
+ /// The minimum index of the column, inclusive.
+ ///
+ [Argument(ArgumentType.Required, HelpText = "First index in the range")]
+ public int Min;
+
+ ///
+ /// The maximum index of the column, inclusive. If
+ /// indicates that the should auto-detect the legnth
+ /// of the lines, and read untill the end.
+ /// If is specified, the field is ignored.
+ ///
+ [Argument(ArgumentType.AtMostOnce, HelpText = "Last index in the range")]
+ public int? Max;
+
+ ///
+ /// Whether this range extends to the end of the line, but should be a fixed number of items.
+ /// If is specified, the field is ignored.
+ ///
+ [Argument(ArgumentType.AtMostOnce,
+ HelpText = "This range extends to the end of the line, but should be a fixed number of items",
+ ShortName = "auto")]
+ public bool AutoEnd;
+
+ ///
+ /// Whether this range includes only other indices not specified.
+ ///
+ [Argument(ArgumentType.AtMostOnce, HelpText = "This range includes only other indices not specified", ShortName = "other")]
+ public bool AllOther;
+
+ ///
+ /// Force scalar columns to be treated as vectors of length one.
+ ///
+ [Argument(ArgumentType.AtMostOnce, HelpText = "Force scalar columns to be treated as vectors of length one", ShortName = "vector")]
+ public bool ForceVector;
+
+ internal static Range FromTextLoaderRange(TextLoader.Range range)
+ {
+ return new Range(range.Min, range.Max);
+ }
+ }
+
///
/// The settings for
///
diff --git a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoaderCatalog.cs b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoaderCatalog.cs
index e1177082d0..011bd39612 100644
--- a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoaderCatalog.cs
+++ b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DatabaseLoaderCatalog.cs
@@ -11,9 +11,7 @@ namespace Microsoft.ML
///
public static class DatabaseLoaderCatalog
{
- ///
- /// Create a database loader .
- ///
+ /// Create a database loader .
/// The catalog.
/// Array of columns defining the schema.
public static DatabaseLoader CreateDatabaseLoader(this DataOperationsCatalog catalog,
@@ -23,8 +21,22 @@ public static DatabaseLoader CreateDatabaseLoader(this DataOperationsCatalog cat
{
Columns = columns,
};
-
- return new DatabaseLoader(CatalogUtils.GetEnvironment(catalog), options);
+ return catalog.CreateDatabaseLoader(options);
}
+
+ /// Create a database loader .
+ /// The catalog.
+ /// Defines the settings of the load operation.
+ public static DatabaseLoader CreateDatabaseLoader(this DataOperationsCatalog catalog,
+ DatabaseLoader.Options options)
+ => new DatabaseLoader(CatalogUtils.GetEnvironment(catalog), options);
+
+ /// Create a database loader .
+ /// Defines the schema of the data to be loaded. Use public fields or properties
+ /// decorated with (and possibly other attributes) to specify the column
+ /// names and their data types in the schema of the loaded data.
+ /// The catalog.
+ public static DatabaseLoader CreateDatabaseLoader(this DataOperationsCatalog catalog)
+ => DatabaseLoader.CreateDatabaseLoader(CatalogUtils.GetEnvironment(catalog));
}
}
diff --git a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DbExtensions.cs b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DbExtensions.cs
index acbf4e1c84..7a01d89232 100644
--- a/src/Microsoft.ML.Experimental/DataLoadSave/Database/DbExtensions.cs
+++ b/src/Microsoft.ML.Experimental/DataLoadSave/Database/DbExtensions.cs
@@ -66,5 +66,82 @@ public static Type ToType(this DbType dbType)
return null;
}
}
+
+ /// Maps a to the associated .
+ public static DbType ToDbType(this InternalDataKind dataKind)
+ {
+ switch (dataKind)
+ {
+ case InternalDataKind.I1:
+ {
+ return DbType.SByte;
+ }
+
+ case InternalDataKind.U1:
+ {
+ return DbType.Byte;
+ }
+
+ case InternalDataKind.I2:
+ {
+ return DbType.Int16;
+ }
+
+ case InternalDataKind.U2:
+ {
+ return DbType.UInt16;
+ }
+
+ case InternalDataKind.I4:
+ {
+ return DbType.Int32;
+ }
+
+ case InternalDataKind.U4:
+ {
+ return DbType.UInt32;
+ }
+
+ case InternalDataKind.I8:
+ {
+ return DbType.Int64;
+ }
+
+ case InternalDataKind.U8:
+ {
+ return DbType.UInt64;
+ }
+
+ case InternalDataKind.R4:
+ {
+ return DbType.Single;
+ }
+
+ case InternalDataKind.R8:
+ {
+ return DbType.Double;
+ }
+
+ case InternalDataKind.TX:
+ {
+ return DbType.String;
+ }
+
+ case InternalDataKind.BL:
+ {
+ return DbType.Boolean;
+ }
+
+ case InternalDataKind.DT:
+ {
+ return DbType.DateTime;
+ }
+
+ default:
+ {
+ throw new NotSupportedException();
+ }
+ }
+ }
}
}
diff --git a/test/Microsoft.ML.Tests/DatabaseLoaderTests.cs b/test/Microsoft.ML.Tests/DatabaseLoaderTests.cs
index c2bc27df75..eb6dc45329 100644
--- a/test/Microsoft.ML.Tests/DatabaseLoaderTests.cs
+++ b/test/Microsoft.ML.Tests/DatabaseLoaderTests.cs
@@ -41,7 +41,7 @@ public void IrisLightGbm()
var loader = mlContext.Data.CreateDatabaseLoader(loaderColumns);
- var mockProviderFactory = new MockProviderFactory(mlContext, loaderColumns);
+ var mockProviderFactory = new MockProviderFactory(mlContext, loader);
var databaseSource = new DatabaseSource(mockProviderFactory, connectionString, commandText);
var trainingData = loader.Load(databaseSource);
@@ -79,18 +79,9 @@ public void IrisSdcaMaximumEntropy()
var connectionString = GetDataPath(TestDatasets.iris.trainFilename);
var commandText = "Label;SepalLength;SepalWidth;PetalLength;PetalWidth";
- var loaderColumns = new DatabaseLoader.Column[]
- {
- new DatabaseLoader.Column() { Name = "Label", Type = DbType.Int32 },
- new DatabaseLoader.Column() { Name = "SepalLength", Type = DbType.Single },
- new DatabaseLoader.Column() { Name = "SepalWidth", Type = DbType.Single },
- new DatabaseLoader.Column() { Name = "PetalLength", Type = DbType.Single },
- new DatabaseLoader.Column() { Name = "PetalWidth", Type = DbType.Single }
- };
-
- var loader = mlContext.Data.CreateDatabaseLoader(loaderColumns);
+ var loader = mlContext.Data.CreateDatabaseLoader();
- var mockProviderFactory = new MockProviderFactory(mlContext, loaderColumns);
+ var mockProviderFactory = new MockProviderFactory(mlContext, loader);
var databaseSource = new DatabaseSource(mockProviderFactory, connectionString, commandText);
var trainingData = loader.Load(databaseSource);
@@ -123,11 +114,15 @@ public void IrisSdcaMaximumEntropy()
public class IrisData
{
+ public int Label;
+
public float SepalLength;
+
public float SepalWidth;
+
public float PetalLength;
+
public float PetalWidth;
- public int Label;
}
public class IrisPrediction
@@ -140,42 +135,39 @@ public class IrisPrediction
internal sealed class MockProviderFactory : DbProviderFactory
{
private MLContext _context;
- private DatabaseLoader.Column[] _columns;
+ private DatabaseLoader _databaseLoader;
- public MockProviderFactory(MLContext context, DatabaseLoader.Column[] columns)
+ public MockProviderFactory(MLContext context, DatabaseLoader databaseLoader)
{
_context = context;
- _columns = columns;
+ _databaseLoader = databaseLoader;
}
- public override DbConnection CreateConnection() => new MockConnection(_context, _columns);
+ public override DbConnection CreateConnection() => new MockConnection(_context, _databaseLoader);
}
internal sealed class MockConnection : DbConnection
{
private string _dataPath;
- private TextLoader _reader;
+ private TextLoader _textLoader;
- public MockConnection(MLContext context, DatabaseLoader.Column[] columns)
+ public MockConnection(MLContext context, DatabaseLoader databaseLoader)
{
- Columns = columns;
-
- var readerColumns = new TextLoader.Column[columns.Length];
+ var outputSchema = databaseLoader.GetOutputSchema();
+ var readerColumns = new TextLoader.Column[outputSchema.Count];
- for (int i = 0; i < columns.Length; i++)
+ for (int i = 0; i < outputSchema.Count; i++)
{
- var column = columns[i];
- var columnType = column.Type.ToType();
+ var column = outputSchema[i];
+ var columnType = column.Type.RawType;
Assert.True(columnType.TryGetDataKind(out var internalDataKind));
readerColumns[i] = new TextLoader.Column(column.Name, internalDataKind.ToDataKind(), i);
}
- _reader = context.Data.CreateTextLoader(readerColumns);
+ _textLoader = context.Data.CreateTextLoader(readerColumns);
}
- public DatabaseLoader.Column[] Columns { get; }
-
public override string ConnectionString
{
get
@@ -205,7 +197,7 @@ public override string ConnectionString
public override void Open()
{
- DataView = _reader.Load(_dataPath);
+ DataView = _textLoader.Load(_dataPath);
}
protected override DbTransaction BeginDbTransaction(IsolationLevel isolationLevel) => throw new NotImplementedException();
@@ -290,7 +282,8 @@ public MockDbDataReader(MockCommand command)
var connection = (MockConnection)_command.Connection;
_dataView = connection.DataView;
- var inputColumns = _dataView.Schema.Where((column) => {
+ var inputColumns = _dataView.Schema.Where((column) =>
+ {
var inputColumnNames = command.CommandText.Split(';');
return inputColumnNames.Any((columnName) => column.Name.Equals(column.Name));
});
@@ -358,19 +351,7 @@ public override int GetInt32(int ordinal)
public override int GetOrdinal(string name)
{
var connection = (MockConnection)_command.Connection;
- var columns = connection.Columns;
-
- for (int i = 0; i < columns.Length; i++)
- {
- var column = columns[i];
-
- if (column.Name.Equals(name))
- {
- return i;
- }
- }
-
- return -1;
+ return connection.DataView.Schema.TryGetColumnIndex(name, out int ordinal) ? ordinal : -1;
}
public override string GetString(int ordinal) => throw new NotImplementedException();