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();