ML TransformLabels

StratifiedTrainTestSplit в Dataset
This commit is contained in:
Mikhalkovich Stanislav 2026-04-01 06:33:26 +03:00
parent ef27d97dc0
commit c8a1b87a2c
6 changed files with 1281 additions and 616 deletions

View file

@ -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.

View file

@ -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

View file

@ -64,7 +64,7 @@ type
ColumnInfo = auto class
Name: string;
ColType: ColumnType;
//IsCategorical: boolean;
//IsCategorical: boolean; - мы убрали это отсюда - только Schema - источник истины!
end;
DataFrameCursor = class;

View file

@ -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;

View file

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