ML - IColumnBoundStep - ограничение для target в Pipeline
This commit is contained in:
parent
431c2a9861
commit
b9b55b513b
|
|
@ -8,10 +8,11 @@ begin
|
|||
|
||||
var target := 'price';
|
||||
|
||||
var (trainDf, testDf) := df.TrainTestSplit(0.2, 42);
|
||||
var (trainDf, testDf) := df.TrainTestSplit(0.2, seed := 42);
|
||||
|
||||
var pipe :=
|
||||
DataPipeline.Build(
|
||||
TaskKind.tkRegression,
|
||||
target,
|
||||
features,
|
||||
new OneHotEncoder('renovation'),
|
||||
|
|
|
|||
|
|
@ -17,12 +17,36 @@ type
|
|||
IMatrixStep = interface(IPipelineStep)
|
||||
end;
|
||||
|
||||
/// Интерфейс шага конвейера, работающего на уровне DataFrame.
|
||||
/// Реализуется препроцессорами, которые выполняют преобразования табличных данных
|
||||
/// до перехода в числовое (Matrix) представление
|
||||
/// Интерфейс шага конвейера, работающего на уровне DataFrame.
|
||||
/// Реализуется препроцессорами, которые выполняют преобразования табличных данных
|
||||
/// до перехода в числовое (Matrix) представление
|
||||
IDataStep = interface(IPipelineStep)
|
||||
end;
|
||||
|
||||
/// Интерфейс шага конвейера, привязанного к одной колонке DataFrame.
|
||||
/// Используется для препроцессоров, выполняющих преобразование конкретной колонки
|
||||
/// (например, кодирование категориальных признаков).
|
||||
/// Позволяет DataPipeline централизованно контролировать операции над колонками
|
||||
/// (в частности, запрещать изменение целевой переменной).
|
||||
IColumnBoundStep = interface
|
||||
/// Имя колонки, к которой применяется преобразование.
|
||||
/// Должно соответствовать существующей колонке DataFrame.
|
||||
/// Используется для валидации (например, запрет преобразования target).
|
||||
property ColumnName: string read;
|
||||
end;
|
||||
|
||||
/// Интерфейс шага конвейера, работающего с несколькими колонками DataFrame.
|
||||
/// Используется для препроцессоров, выполняющих преобразования над набором колонок
|
||||
/// (например, заполнение пропусков, масштабирование или кодирование нескольких признаков).
|
||||
/// Позволяет DataPipeline контролировать операции над колонками
|
||||
/// (в частности, предотвращать изменение целевой переменной).
|
||||
IColumnsBoundStep = interface
|
||||
/// Список колонок, к которым применяется преобразование.
|
||||
/// Каждое имя должно соответствовать существующей колонке DataFrame.
|
||||
/// Используется для валидации (например, запрет преобразования target).
|
||||
property Columns: array of string read;
|
||||
end;
|
||||
|
||||
/// Базовый интерфейс модели машинного обучения
|
||||
IModel = interface(IMatrixStep)
|
||||
function Predict(X: Matrix): Vector;
|
||||
|
|
|
|||
|
|
@ -3937,6 +3937,8 @@ begin
|
|||
Log2Features: Result := integer(Log2(p));
|
||||
HalfFeatures: Result := p div 2;
|
||||
end;
|
||||
if Result < 1 then
|
||||
Result := 1;
|
||||
end;
|
||||
|
||||
procedure RandomForestBase.BootstrapRowIndices(n: integer; var rows: array of integer);
|
||||
|
|
|
|||
|
|
@ -350,6 +350,11 @@ const
|
|||
'LabelEncoder нельзя применять к целевому столбцу. Используйте EncodeLabels!!LabelEncoder cannot be applied to the target column. Use EncodeLabels instead';
|
||||
ER_ENCODELABELS_NOT_CATEGORICAL =
|
||||
'Целевой столбец должен быть категориальным для задач классификации!!Target column must be categorical for classification tasks';
|
||||
ER_PIPELINE_TARGET_TRANSFORM_NOT_ALLOWED =
|
||||
'Преобразование целевой переменной "{0}" запрещено в DataPipeline!!' +
|
||||
'Transformation of target variable "{0}" is not allowed in DataPipeline';
|
||||
|
||||
|
||||
//-----------------------------
|
||||
// DataPipeline
|
||||
//-----------------------------
|
||||
|
|
@ -374,11 +379,17 @@ begin
|
|||
Error(ER_PIPELINE_MODIFY_AFTER_FIT);
|
||||
|
||||
// --- target protection
|
||||
// NOTE: Only built-in LabelEncoder is restricted for target.
|
||||
// Custom preprocessors are not checked.
|
||||
if step is LabelEncoder(var enc) then
|
||||
if enc.ColumnName = fTarget then
|
||||
ArgumentError(ER_LABELENCODER_TARGET_NOT_ALLOWED, fTarget);
|
||||
// Любой шаг DataFrame, привязанный к одной или нескольким колонкам,
|
||||
// не должен затрагивать целевую переменную (target).
|
||||
// Проверка выполняется через интерфейсы IColumnBoundStep / IColumnsBoundStep
|
||||
// без привязки к конкретным классам.
|
||||
if step is IColumnBoundStep(var cstep) then
|
||||
if cstep.ColumnName = fTarget then
|
||||
ArgumentError(ER_PIPELINE_TARGET_TRANSFORM_NOT_ALLOWED, fTarget);
|
||||
|
||||
if step is IColumnsBoundStep(var mstep) then
|
||||
if fTarget in mstep.Columns then
|
||||
ArgumentError(ER_PIPELINE_TARGET_TRANSFORM_NOT_ALLOWED, fTarget);
|
||||
|
||||
// --- DataFrame step
|
||||
if step is IPreprocessor then
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ type
|
|||
/// в порядке первого появления категорий.
|
||||
/// Работает только со строковыми столбцами и предназначен для признаков.
|
||||
/// Не должен применяться к целевому столбцу (target).
|
||||
LabelEncoder = class(IPreprocessor)
|
||||
LabelEncoder = class(IPreprocessor, IColumnBoundStep)
|
||||
private
|
||||
col: string;
|
||||
mapping: Dictionary<string, integer>;
|
||||
|
|
@ -68,7 +68,7 @@ type
|
|||
/// Категории фиксируются при Fit
|
||||
/// Неизвестные категории кодируются нулями
|
||||
/// Пропущенные значения (NA) кодируются нулями
|
||||
OneHotEncoder = class(IPreprocessor)
|
||||
OneHotEncoder = class(IPreprocessor, IColumnBoundStep)
|
||||
private
|
||||
col: string;
|
||||
categories: array of string;
|
||||
|
|
@ -87,6 +87,8 @@ type
|
|||
function FitTransform(df: DataFrame): DataFrame;
|
||||
|
||||
function ToString: string; override;
|
||||
|
||||
property ColumnName: string read col;
|
||||
end;
|
||||
|
||||
ImputeStrategy = (isMean, isConstant);
|
||||
|
|
@ -94,7 +96,7 @@ type
|
|||
/// Заполняет пропущенные значения (NA) в числовых столбцах
|
||||
/// Поддерживает стратегии isMean и isConstant
|
||||
/// Работает только с Int и Float столбцами
|
||||
Imputer = class(IPreprocessor)
|
||||
Imputer = class(IPreprocessor, IColumnsBoundStep)
|
||||
private
|
||||
cols: array of string;
|
||||
strategy: ImputeStrategy;
|
||||
|
|
@ -118,6 +120,8 @@ type
|
|||
function FitTransform(df: DataFrame): DataFrame;
|
||||
|
||||
function ToString: string; override;
|
||||
|
||||
property Columns: array of string read cols;
|
||||
end;
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue