ML TransformLabels
StratifiedTrainTestSplit в Dataset
This commit is contained in:
parent
ef27d97dc0
commit
c8a1b87a2c
|
|
@ -0,0 +1,25 @@
|
|||
uses MLABC;
|
||||
|
||||
begin
|
||||
var ds := Datasets.Iris;
|
||||
|
||||
// --- stratified split
|
||||
var (train, test) := ds.StratifiedTrainTestSplit(0.2, 42);
|
||||
|
||||
var Xtrain := train.Data.ToMatrix(train.Features);
|
||||
var Xtest := test.Data.ToMatrix(test.Features);
|
||||
|
||||
// неудобно, но правильно
|
||||
var classes: array of string;
|
||||
|
||||
var ytrain := train.Data.EncodeLabels(train.Target, classes);
|
||||
var ytest := test.Data.TransformLabels(test.Target, classes);
|
||||
|
||||
// --- модель
|
||||
var model := new LogisticRegression;
|
||||
model.Fit(Xtrain, ytrain);
|
||||
|
||||
var pred := model.Predict(Xtest);
|
||||
|
||||
Println('Accuracy:', Metrics.Accuracy(ytest, pred):0:3);
|
||||
end.
|
||||
|
|
@ -25,6 +25,13 @@ function EncodeLabels(labels: array of string): array of integer;
|
|||
/// Используется при обучении моделей и визуализации
|
||||
function EncodeLabels(labels: array of string; var classes: array of string): array of integer;
|
||||
|
||||
/// Преобразует строковые метки классов в целочисленные индексы
|
||||
/// с использованием заранее заданного массива classes (mapping).
|
||||
/// classes должен быть получен из EncodeLabels.
|
||||
/// Если встречается неизвестная метка — выбрасывается исключение.
|
||||
/// Используется для применения кодирования к тестовым данным (Transform).
|
||||
function TransformLabels(labels: array of string; classes: array of string): array of integer;
|
||||
|
||||
/// Преобразует целочисленные индексы классов обратно в строковые метки.
|
||||
/// Массив classes задаёт соответствие: classes[i] — имя класса с индексом i.
|
||||
/// Используется для получения текстовых предсказаний моделей.
|
||||
|
|
@ -53,6 +60,8 @@ const
|
|||
'Столбец "{0}" должен быть категориальным для EncodeLabels!!Column "{0}" must be categorical for EncodeLabels';
|
||||
ER_ENCODELABELS_UNSUPPORTED_TYPE =
|
||||
'Неподдерживаемый тип столбца "{0}" для EncodeLabels!!Unsupported column type "{0}" for EncodeLabels';
|
||||
ER_UNKNOWN_CLASS_IN_TRANSFORM =
|
||||
'Неизвестное значение класса "{0}" при преобразовании меток!!Unknown class value "{0}" in TransformLabels';
|
||||
|
||||
function LabelsToInts(y: Vector): array of integer;
|
||||
begin
|
||||
|
|
@ -90,6 +99,30 @@ begin
|
|||
Result := res;
|
||||
end;
|
||||
|
||||
function TransformLabels(labels: array of string; classes: array of string): array of integer;
|
||||
begin
|
||||
if labels = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'labels');
|
||||
|
||||
var map := new Dictionary<string, integer>;
|
||||
for var i := 0 to classes.Length - 1 do
|
||||
map[classes[i]] := i;
|
||||
|
||||
var res := new integer[labels.Length];
|
||||
|
||||
for var i := 0 to labels.Length - 1 do
|
||||
begin
|
||||
var lbl := labels[i];
|
||||
|
||||
if not map.ContainsKey(lbl) then
|
||||
Error('Unknown class in test');
|
||||
|
||||
res[i] := map[lbl];
|
||||
end;
|
||||
|
||||
Result := res;
|
||||
end;
|
||||
|
||||
function EncodeLabels(labels: array of string): array of integer;
|
||||
begin
|
||||
var classes: array of string;
|
||||
|
|
@ -273,6 +306,121 @@ begin
|
|||
end;
|
||||
end;
|
||||
|
||||
/// Преобразует строковые метки целевого столбца в целочисленные индексы (0,1,2,...)
|
||||
/// с использованием заданного массива classes (mapping индекс → метка).
|
||||
/// classes должен быть получен ранее с помощью EncodeLabels.
|
||||
/// При обнаружении неизвестного значения выбрасывается исключение.
|
||||
/// Используется для применения кодирования к тестовым данным (Transform).
|
||||
function TransformLabels(Self: DataFrame; target: string; classes: array of string): array of integer; extensionmethod;
|
||||
begin
|
||||
if Self = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'Self');
|
||||
|
||||
if target = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'target');
|
||||
|
||||
if classes = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'classes');
|
||||
|
||||
if not Self.HasColumn(target) then
|
||||
ArgumentError(ER_COLUMN_NOT_FOUND, target);
|
||||
|
||||
if not Self.IsCategorical(target) then
|
||||
ArgumentError(ER_ENCODELABELS_NOT_CATEGORICAL, target);
|
||||
|
||||
// --- строим mapping
|
||||
var map := new Dictionary<string, integer>;
|
||||
for var i := 0 to classes.Length - 1 do
|
||||
map[classes[i]] := i;
|
||||
|
||||
case Self.GetColumnType(target) of
|
||||
|
||||
ColumnType.ctStr:
|
||||
begin
|
||||
var data := Self.GetStrColumn(target).ToArray;
|
||||
var res := new integer[data.Length];
|
||||
|
||||
for var i := 0 to data.Length - 1 do
|
||||
begin
|
||||
var lbl := data[i];
|
||||
|
||||
if not map.ContainsKey(lbl) then
|
||||
Error(ER_UNKNOWN_CLASS_IN_TRANSFORM, lbl);
|
||||
|
||||
res[i] := map[lbl];
|
||||
end;
|
||||
|
||||
Result := res;
|
||||
end;
|
||||
|
||||
ColumnType.ctInt:
|
||||
begin
|
||||
var data := Self.GetIntColumn(target).ToArray;
|
||||
var res := new integer[data.Length];
|
||||
|
||||
for var i := 0 to data.Length - 1 do
|
||||
begin
|
||||
var lbl := data[i].ToString;
|
||||
|
||||
if not map.ContainsKey(lbl) then
|
||||
Error('Unknown class in TransformLabels');
|
||||
|
||||
res[i] := map[lbl];
|
||||
end;
|
||||
|
||||
Result := res;
|
||||
end;
|
||||
|
||||
else
|
||||
ArgumentError(ER_ENCODELABELS_UNSUPPORTED_TYPE, target);
|
||||
end;
|
||||
end;
|
||||
|
||||
/// Преобразует значения целочисленного категориального столбца
|
||||
/// в плотные индексы (0,1,2,...) с использованием заданного массива classes.
|
||||
/// classes должен быть получен ранее с помощью EncodeLabelsInt.
|
||||
/// При обнаружении неизвестного значения выбрасывается исключение.
|
||||
function TransformLabelsInt(Self: DataFrame; target: string; classes: array of integer): array of integer; extensionmethod;
|
||||
begin
|
||||
if Self = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'Self');
|
||||
|
||||
if target = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'target');
|
||||
|
||||
if classes = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'classes');
|
||||
|
||||
if not Self.HasColumn(target) then
|
||||
ArgumentError(ER_COLUMN_NOT_FOUND, target);
|
||||
|
||||
if not Self.IsCategorical(target) then
|
||||
ArgumentError(ER_ENCODELABELS_NOT_CATEGORICAL, target);
|
||||
|
||||
if Self.GetColumnType(target) <> ColumnType.ctInt then
|
||||
ArgumentError(ER_ENCODELABELS_UNSUPPORTED_TYPE, target);
|
||||
|
||||
// --- mapping: значение → индекс
|
||||
var map := new Dictionary<integer, integer>;
|
||||
for var i := 0 to classes.Length - 1 do
|
||||
map[classes[i]] := i;
|
||||
|
||||
var data := Self.GetIntColumn(target).ToArray;
|
||||
var res := new integer[data.Length];
|
||||
|
||||
for var i := 0 to data.Length - 1 do
|
||||
begin
|
||||
var v := data[i];
|
||||
|
||||
if not map.ContainsKey(v) then
|
||||
Error(ER_UNKNOWN_CLASS_IN_TRANSFORM, v);
|
||||
|
||||
res[i] := map[v];
|
||||
end;
|
||||
|
||||
Result := res;
|
||||
end;
|
||||
|
||||
/// Кодирует значения целочисленного категориального столбца DataFrame
|
||||
/// в плотные целочисленные индексы 0,1,2,...
|
||||
/// Каждому уникальному значению присваивается номер в порядке первого появления.
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -64,7 +64,7 @@ type
|
|||
ColumnInfo = auto class
|
||||
Name: string;
|
||||
ColType: ColumnType;
|
||||
//IsCategorical: boolean;
|
||||
//IsCategorical: boolean; - мы убрали это отсюда - только Schema - источник истины!
|
||||
end;
|
||||
|
||||
DataFrameCursor = class;
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ type
|
|||
Statistics = DataFrameABC.Statistics;
|
||||
CsvLoader = DataFrameABC.CsvLoader;
|
||||
JoinKind = DataFrameABC.JoinKind;
|
||||
GroupView = DataFrameABC.GroupView;
|
||||
|
||||
IProbabilisticClassifier = MLCoreABC.IProbabilisticClassifier;
|
||||
IRegressor = MLCoreABC.IRegressor;
|
||||
|
|
@ -87,6 +88,7 @@ type
|
|||
Imputer = PreprocessorABC.Imputer;
|
||||
|
||||
Datasets = MLDatasets.Datasets;
|
||||
Dataset = MLDatasets.Dataset;
|
||||
|
||||
IModel = MLCoreABC.IModel;
|
||||
ISupervisedModel = MLCoreABC.ISupervisedModel;
|
||||
|
|
@ -95,6 +97,16 @@ type
|
|||
UDataPipeline = MLPipelineABC.UDataPipeline;
|
||||
TaskKind = MLPipelineABC.TaskKind;
|
||||
|
||||
AggregationKind = DataFrameABC.AggregationKind;
|
||||
|
||||
const
|
||||
akMean = AggregationKind.akMean;
|
||||
akMin = AggregationKind.akMin;
|
||||
akMax = AggregationKind.akMax;
|
||||
akCount = AggregationKind.akCount;
|
||||
akSum = AggregationKind.akSum;
|
||||
akStd = AggregationKind.akStd;
|
||||
|
||||
function LabelsToInts(y: Vector): array of integer;
|
||||
function EncodeLabels(labels: array of string): array of integer;
|
||||
|
||||
|
|
|
|||
|
|
@ -37,8 +37,18 @@ type
|
|||
function IsSupervised: boolean;
|
||||
|
||||
/// Разбивает датасет на обучающую и тестовую части.
|
||||
/// testRatio — доля тестовой выборки.
|
||||
function TrainTestSplit(testRatio: real := 0.2; seed: integer := -1): (Dataset, Dataset);
|
||||
/// testRatio — доля тестовой выборки (0 < testRatio < 1).
|
||||
/// shuffle — перемешивать ли строки перед разбиением.
|
||||
/// seed — начальное значение генератора случайных чисел (для воспроизводимости).
|
||||
/// Возвращает два датасета: (train, test).
|
||||
function TrainTestSplit(testRatio: real := 0.2; shuffle: boolean := True; seed: integer := -1): (Dataset, Dataset);
|
||||
|
||||
/// Выполняет стратифицированное разбиение датасета на обучающую и тестовую части.
|
||||
/// Сохраняет распределение значений целевой переменной в обеих выборках.
|
||||
/// testRatio — доля тестовой выборки (0 < testRatio < 1).
|
||||
/// seed — начальное значение генератора случайных чисел (для воспроизводимости).
|
||||
/// Возвращает два датасета: (train, test).
|
||||
function StratifiedTrainTestSplit(testRatio: real := 0.2; seed: integer := -1): (Dataset, Dataset);
|
||||
|
||||
/// Возвращает первые n строк таблицы данных.
|
||||
function Head(n: integer := 5): DataFrame;
|
||||
|
|
@ -314,7 +324,13 @@ const
|
|||
'Параметр {0} должен быть в диапазоне (0, 1)!!Parameter {0} must be in range (0, 1)';
|
||||
ER_PARAM_LT =
|
||||
'Параметр {0} должен быть меньше допустимого значения!!Parameter {0} must be less than allowed value';
|
||||
|
||||
ER_TEST_RATIO_INVALID =
|
||||
'Некорректное значение testRatio (должно быть между 0 и 1)!!Invalid testRatio (must be between 0 and 1)';
|
||||
ER_GROUPBY_UNSUPPORTED_KEY_TYPE =
|
||||
'Неподдерживаемый тип ключа для группировки!!Unsupported key type for grouping';
|
||||
ER_STRATIFIED_ONLY_FOR_CLASSIFICATION =
|
||||
'Стратифицированное разбиение доступно только для задач классификации!!Stratified split is only for classification tasks';
|
||||
|
||||
C_DATASET = 'Датасет: {0}!!Dataset: {0}';
|
||||
C_DESCRIPTION = 'Описание:!!Description:';
|
||||
C_TASK = 'Задача: {0}!!Task: {0}';
|
||||
|
|
@ -373,18 +389,120 @@ begin
|
|||
Result.Name := Name;
|
||||
Result.Data := df;
|
||||
|
||||
Result.Features := Features;
|
||||
Result.Features := Copy(Features);
|
||||
Result.Target := Target;
|
||||
Result.Task := Task;
|
||||
|
||||
Result.Description := Description;
|
||||
Result.FeatureLabels := FeatureLabels;
|
||||
Result.ValueLabels := ValueLabels;
|
||||
|
||||
Result.FeatureLabels := new Dictionary<string,string>(FeatureLabels);
|
||||
|
||||
Result.ValueLabels := new Dictionary<string, Dictionary<string,string>>;
|
||||
foreach var kvp in ValueLabels do
|
||||
Result.ValueLabels[kvp.Key] := new Dictionary<string,string>(kvp.Value);
|
||||
end;
|
||||
|
||||
function Dataset.TrainTestSplit(testRatio: real; seed: integer): (Dataset, Dataset);
|
||||
function Dataset.TrainTestSplit(testRatio: real; shuffle: boolean; seed: integer): (Dataset, Dataset);
|
||||
begin
|
||||
var (trainDf, testDf) := Data.TrainTestSplit(testRatio, seed);
|
||||
var (trainDf, testDf) := Data.TrainTestSplit(testRatio, shuffle, seed);
|
||||
|
||||
var trainDs := CloneMeta(trainDf);
|
||||
var testDs := CloneMeta(testDf);
|
||||
|
||||
Result := (trainDs, testDs);
|
||||
end;
|
||||
|
||||
function Dataset.StratifiedTrainTestSplit(testRatio: real; seed: integer): (Dataset, Dataset);
|
||||
begin
|
||||
if Data = nil then
|
||||
ArgumentNullError(ER_ARG_NULL, 'Data');
|
||||
|
||||
if (testRatio <= 0.0) or (testRatio >= 1.0) then
|
||||
ArgumentError(ER_TEST_RATIO_INVALID, testRatio);
|
||||
|
||||
if Task <> Classification then
|
||||
Error(ER_STRATIFIED_ONLY_FOR_CLASSIFICATION);
|
||||
|
||||
var n := Data.RowCount;
|
||||
if n < 2 then
|
||||
ArgumentError(ER_EMPTY_DATA, 'StratifiedTrainTestSplit');
|
||||
|
||||
var actualSeed := if seed >= 0 then seed else System.Environment.TickCount and integer.MaxValue;
|
||||
var rnd := new System.Random(actualSeed);
|
||||
|
||||
// --- target column
|
||||
var ci := Data.ColumnIndex(Target);
|
||||
var col := Data.GetColumn(ci);
|
||||
|
||||
var groups := new Dictionary<object, List<integer>>;
|
||||
|
||||
// --- группировка по target
|
||||
case col.Info.ColType of
|
||||
|
||||
ctInt:
|
||||
begin
|
||||
var data := IntColumn(col).Data;
|
||||
|
||||
for var i := 0 to n - 1 do
|
||||
begin
|
||||
var key: object := data[i];
|
||||
|
||||
var lst: List<integer>;
|
||||
if not groups.TryGetValue(key, lst) then
|
||||
begin
|
||||
lst := new List<integer>;
|
||||
groups[key] := lst;
|
||||
end;
|
||||
|
||||
lst.Add(i);
|
||||
end;
|
||||
end;
|
||||
|
||||
ctStr:
|
||||
begin
|
||||
var data := StrColumn(col).Data;
|
||||
|
||||
for var i := 0 to n - 1 do
|
||||
begin
|
||||
var key: object := data[i];
|
||||
|
||||
var lst: List<integer>;
|
||||
if not groups.TryGetValue(key, lst) then
|
||||
begin
|
||||
lst := new List<integer>;
|
||||
groups[key] := lst;
|
||||
end;
|
||||
|
||||
lst.Add(i);
|
||||
end;
|
||||
end;
|
||||
|
||||
else
|
||||
Error(ER_GROUPBY_UNSUPPORTED_KEY_TYPE, col.Info.ColType);
|
||||
end;
|
||||
|
||||
var trainIdx := new List<integer>;
|
||||
var testIdx := new List<integer>;
|
||||
|
||||
// --- split внутри каждой группы
|
||||
foreach var kvp in groups do
|
||||
begin
|
||||
var arr := kvp.Value.ToArray;
|
||||
arr.Shuffle(rnd);
|
||||
|
||||
var m := arr.Length;
|
||||
var rawSize := Round(m * testRatio);
|
||||
var testSize := rawSize.Clamp(1, m - 1);
|
||||
|
||||
for var i := 0 to testSize - 1 do
|
||||
testIdx.Add(arr[i]);
|
||||
|
||||
for var i := testSize to m - 1 do
|
||||
trainIdx.Add(arr[i]);
|
||||
end;
|
||||
|
||||
var trainDf := Data.TakeRows(trainIdx.ToArray);
|
||||
var testDf := Data.TakeRows(testIdx.ToArray);
|
||||
|
||||
var trainDs := CloneMeta(trainDf);
|
||||
var testDs := CloneMeta(testDf);
|
||||
|
|
|
|||
Loading…
Reference in a new issue