From 947fd12b3af0122a75e8d6257e59f9366c5f56b8 Mon Sep 17 00:00:00 2001 From: obroccolio Date: Fri, 11 Sep 2026 11:27:49 +0800 Subject: [PATCH] feat: split training targets by task type --- README.md | 1 + src/lmpm/training/targets.py | 81 +++++++++++++++++++++++++++++++++ tests/test_training_targets.py | 82 ++++++++++++++++++++++++++++++++++ 3 files changed, 164 insertions(+) create mode 100644 src/lmpm/training/targets.py create mode 100644 tests/test_training_targets.py diff --git a/README.md b/README.md index 5a77257..300d542 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,7 @@ tech-architecture/ 技术架构报告与本地 Mermaid 资源 - `ExperimentRecord`:加工参数、质量指标与记录时间。 - SQLite 骨架:`materials` 与 `experiments` 两张核心表,支持外键约束与实验记录往返。 - 训练数据接口:从外部数据平台的 `material_records` 表只读加载记录,并分离元信息、特征和目标。 +- 目标拆分接口:按分类/回归任务组织 `is_cut_through`、`etching_depth` 等实验结果。 - 工程配置:`pyproject.toml` 统一依赖、pytest 与 Ruff 配置。 后续将按架构报告逐步补齐 DoE 生成、CSV 交换、设备采集、建模管线和推理服务。 diff --git a/src/lmpm/training/targets.py b/src/lmpm/training/targets.py new file mode 100644 index 0000000..6538a84 --- /dev/null +++ b/src/lmpm/training/targets.py @@ -0,0 +1,81 @@ +from typing import Literal + +from pydantic import BaseModel + +from lmpm.training.dataset import TrainingDataset + +TargetTask = Literal["classification", "regression"] + + +class TargetSpec(BaseModel): + """The learning-task type of one platform outcome field.""" + + name: str + task: TargetTask + + +class TargetSplit(BaseModel): + """Outcome values grouped by compatible learning task.""" + + classification: dict[str, tuple[bool, ...]] + regression: dict[str, tuple[float, ...]] + + +TARGET_SPECS = ( + TargetSpec(name="is_cut_through", task="classification"), + TargetSpec(name="carbonized_edge_width", task="regression"), + TargetSpec(name="etching_depth", task="regression"), + TargetSpec(name="is_fire_smolder", task="classification"), + TargetSpec(name="pattern_clarity_score", task="regression"), + TargetSpec(name="presentation_balance_score", task="regression"), +) + + +def target_specs() -> tuple[TargetSpec, ...]: + """Return the stable mapping from platform outcomes to task types.""" + return TARGET_SPECS + + +def split_targets(dataset: TrainingDataset) -> TargetSplit: + """Split complete outcome columns by learning task. + + Empty datasets are allowed so dataset profiling can run before data + collection starts. Any incomplete target value is an error because a + training row must provide every outcome it is being trained against. + """ + classification: dict[str, list[bool]] = {} + regression: dict[str, list[float]] = {} + + for spec in TARGET_SPECS: + values: list[bool | float] = [] + for record in dataset.records: + if spec.name not in record.targets or record.targets[spec.name] is None: + raise ValueError( + f"{record.metadata['experiment_id']}: missing target {spec.name}" + ) + + value = record.targets[spec.name] + if spec.task == "classification": + if not isinstance(value, bool): + raise ValueError( + f"{record.metadata['experiment_id']}: " + f"{spec.name} must be boolean" + ) + values.append(value) + else: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError( + f"{record.metadata['experiment_id']}: " + f"{spec.name} must be numeric" + ) + values.append(float(value)) + + if spec.task == "classification": + classification[spec.name] = values + else: + regression[spec.name] = values + + return TargetSplit( + classification={name: tuple(values) for name, values in classification.items()}, + regression={name: tuple(values) for name, values in regression.items()}, + ) diff --git a/tests/test_training_targets.py b/tests/test_training_targets.py new file mode 100644 index 0000000..44c39f2 --- /dev/null +++ b/tests/test_training_targets.py @@ -0,0 +1,82 @@ +import pytest +from pydantic import ValidationError + +from lmpm.training.dataset import TrainingDataset, TrainingRecord +from lmpm.training.targets import TargetSpec, split_targets, target_specs + + +def make_record(number: int = 1, *, target_overrides=None) -> TrainingRecord: + targets = { + "is_cut_through": True, + "carbonized_edge_width": 0.2, + "etching_depth": 30.0, + "is_fire_smolder": False, + "pattern_clarity_score": 9.0, + "presentation_balance_score": 8.0, + } + if target_overrides: + targets.update(target_overrides) + + return TrainingRecord( + metadata={"experiment_id": f"LAS-2026-{number:04d}"}, + features={"actual_output_power": 8.5}, + targets=targets, + ) + + +def test_target_specs_describe_all_platform_targets(): + specs = target_specs() + + assert [(spec.name, spec.task) for spec in specs] == [ + ("is_cut_through", "classification"), + ("carbonized_edge_width", "regression"), + ("etching_depth", "regression"), + ("is_fire_smolder", "classification"), + ("pattern_clarity_score", "regression"), + ("presentation_balance_score", "regression"), + ] + + +def test_target_spec_rejects_unknown_task_type(): + with pytest.raises(ValidationError): + TargetSpec(name="unsupported_target", task="unknown") + + +def test_split_targets_groups_by_task_type(): + dataset = TrainingDataset( + records=( + make_record(1), + make_record(2, target_overrides={"is_cut_through": False}), + ) + ) + + split = split_targets(dataset) + + assert split.classification["is_cut_through"] == (True, False) + assert split.classification["is_fire_smolder"] == (False, False) + assert split.regression["carbonized_edge_width"] == pytest.approx((0.2, 0.2)) + assert split.regression["etching_depth"] == pytest.approx((30.0, 30.0)) + assert split.regression["pattern_clarity_score"] == pytest.approx((9.0, 9.0)) + assert split.regression["presentation_balance_score"] == pytest.approx((8.0, 8.0)) + + +def test_empty_dataset_has_empty_target_groups(): + split = split_targets(TrainingDataset(records=())) + + assert split.classification == {"is_cut_through": (), "is_fire_smolder": ()} + assert split.regression == { + "carbonized_edge_width": (), + "etching_depth": (), + "pattern_clarity_score": (), + "presentation_balance_score": (), + } + + +@pytest.mark.parametrize("target_name", ["is_cut_through", "etching_depth"]) +def test_missing_target_value_is_rejected(target_name): + dataset = TrainingDataset( + records=(make_record(target_overrides={target_name: None}),) + ) + + with pytest.raises(ValueError, match=target_name): + split_targets(dataset)