| """승인 P-track paired source로 0.6 raster decoder만 보존형 미세조정한다.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import random |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from torch.utils.data import DataLoader, WeightedRandomSampler |
|
|
| PROJECT_ROOT = Path(__file__).parents[1] |
| SOURCE_ROOT = PROJECT_ROOT / "src" |
| if str(SOURCE_ROOT) not in sys.path: |
| sys.path.insert(0, str(SOURCE_ROOT)) |
| if str(PROJECT_ROOT / "scripts") not in sys.path: |
| sys.path.insert(0, str(PROJECT_ROOT / "scripts")) |
|
|
| from math_grid_drawer.research.ink06_federation import ( |
| FederatedPairedInk06Dataset, federation_provenance06, interpolate_state_dict06, load_product_federation06, |
| resolve_training_device06, source_label_balanced_sampler06, |
| ) |
| from math_grid_drawer.research.math_ink_06 import MathInk06Engine, MathInk06Model |
| from math_grid_drawer.research.math_ink_06 import virtual_raster_similarity06 |
| from train_math_ink_06_candidate import _losses |
| from train_math_ink_06_federated_online import _evaluate_sources, _partition, _source_subset |
|
|
|
|
| def _macro_raster_score(metrics: dict[str, dict[str, float]]) -> float: |
| """필요 변수: source별 raster 지표. 작동 원리: top-1을 우선하고 top-5를 보조하는 source-macro 선택 점수를 만든다.""" |
|
|
| return float(np.mean([row["raster_top1"] + 0.25 * row["raster_top5"] for row in metrics.values()])) |
|
|
|
|
| def _evaluate_geometry_sources( |
| engine: MathInk06Engine, groups: dict[str, list[dict]], exact_to_index: dict[str, int], |
| family_to_index: dict[str, int], batch_size: int, |
| ) -> dict[str, dict[str, float]]: |
| """필요 변수: source holdout·vectorizer. 작동 원리: label과 무관한 score-top1·top-4 raster 재구성도를 계산한다.""" |
|
|
| reports = {} |
| engine.model.eval() |
| for source_id, records in groups.items(): |
| loader = DataLoader( |
| FederatedPairedInk06Dataset(records, exact_to_index, family_to_index), |
| batch_size=batch_size, shuffle=False, num_workers=0, |
| ) |
| selected_total = oracle_total = samples = 0 |
| with torch.inference_mode(): |
| for _online, raster, _coordinates, _states, _target, _family, _source in loader: |
| raster = raster.to(engine.device) |
| output = engine.model.forward_raster(raster) |
| similarity = virtual_raster_similarity06( |
| output["coordinates"], raster, state_logits=output["state_logits"], size=32, sigma=0.025, |
| ) |
| selected = output["hypothesis_scores"].argmax(dim=1) |
| batch_index = torch.arange(len(raster), device=engine.device) |
| selected_total += float(similarity[batch_index, selected].sum()) |
| oracle_total += float(similarity.amax(dim=1).sum()) |
| samples += len(raster) |
| reports[source_id] = { |
| "samples": samples, "score_top1_similarity": selected_total / max(samples, 1), |
| "top4_oracle_similarity": oracle_total / max(samples, 1), |
| } |
| return reports |
|
|
|
|
| def _macro_geometry_score(metrics: dict[str, dict[str, float]]) -> float: |
| """필요 변수: source별 재구성도. 작동 원리: 표본 수 편향 없이 top-4 기하 상한을 평균한다.""" |
|
|
| return float(np.mean([row["top4_oracle_similarity"] for row in metrics.values()])) |
|
|
|
|
| def _hard_label_sampler(records: list[dict], *, seed: int, hard_labels: set[str], multiplier: float) -> WeightedRandomSampler: |
| """필요 변수: source-balanced record·hard label. 작동 원리: 기존 source/label 균형 위에서 공통 실패 기호만 제한적으로 재표집한다.""" |
|
|
| base = source_label_balanced_sampler06(records, seed=seed, samples=len(records)) |
| weights = base.weights.detach().clone() |
| if multiplier < 1.0: |
| raise ValueError("hard label multiplier는 1 이상이어야 합니다.") |
| for index, record in enumerate(records): |
| if str(record["label"]) in hard_labels: |
| weights[index] *= multiplier |
| generator = torch.Generator().manual_seed(seed) |
| return WeightedRandomSampler(weights, len(records), replacement=True, generator=generator) |
|
|
|
|
| def _split_sources(sources, *, seed: int, train_max: int, validation_max: int, test_max: int): |
| """필요 변수: 승인 source·subset 상한. 작동 원리: 기존 federation과 동일한 writer/origin 분리로 train·validation·test를 고정한다.""" |
|
|
| training: list[dict] = [] |
| validation: dict[str, list[dict]] = {} |
| test: dict[str, list[dict]] = {} |
| for source_index, source in enumerate(sources): |
| eligible = [row for row in source.records if row.get("eligible_for_training")] |
| explicit_validation = [row for row in source.records if str(row.get("split")) in {"validation", "valid", "val"}] |
| train_candidates = eligible if explicit_validation else [row for row in eligible if _partition(row) >= 2] |
| validation_candidates = explicit_validation or [row for row in eligible if _partition(row) == 0] |
| training.extend(_source_subset(train_candidates, train_max, seed + source_index)) |
| validation[source.source_id] = _source_subset(validation_candidates, validation_max, seed + 20 + source_index) |
| test_candidates = [row for row in source.records if row.get("split") == "test"] |
| test[source.source_id] = _source_subset(test_candidates, test_max, seed + 40 + source_index) |
| if any(not rows for rows in validation.values()) or any(not rows for rows in test.values()): |
| raise ValueError("source validation/test partition이 비었습니다.") |
| return training, validation, test |
|
|
|
|
| def main() -> None: |
| """필요 변수: 0.6 checkpoint·승인 federation. 작동 원리: online head를 고정하고 hard-label decoder 후보를 holdout으로 선택한다.""" |
|
|
| parser = argparse.ArgumentParser(description="Train Math Ink 0.6 federated raster decoder") |
| parser.add_argument("--checkpoint", type=Path, required=True) |
| parser.add_argument("--registry", type=Path, default=PROJECT_ROOT / "research/dataset_registry.json") |
| parser.add_argument("--source-registry", type=Path, default=PROJECT_ROOT / "research/math_ink_06_source_registry.json") |
| parser.add_argument("--commercial", type=Path, default=PROJECT_ROOT / "research/data/external_trajectory_v1/commercial_ccby4.jsonl.gz") |
| parser.add_argument("--hwrt", type=Path, default=PROJECT_ROOT / "research/data/open_pretrain/hwrt_expanded_v2/hwrt_expanded.jsonl.gz") |
| parser.add_argument("--approval", type=Path, default=PROJECT_ROOT / "research/approvals/HWRT-ODBL-USE-APPROVAL-v1.json") |
| parser.add_argument("--output", type=Path, required=True) |
| parser.add_argument("--hard-labels", default="2,%,A,\\Delta,\\Leftrightarrow,\\mathbb{H},\\mu,\\varpi,p") |
| parser.add_argument("--hard-label-multiplier", type=float, default=2.0) |
| parser.add_argument("--max-train-per-source", type=int, default=1000) |
| parser.add_argument("--max-validation-per-source", type=int, default=300) |
| parser.add_argument("--max-test-per-source", type=int, default=500) |
| parser.add_argument("--batch-size", type=int, default=32) |
| parser.add_argument("--epochs", type=int, default=1) |
| parser.add_argument("--learning-rate", type=float, default=1e-5) |
| parser.add_argument("--cycle-weight", type=float, default=0.05) |
| parser.add_argument("--preservation-weight", type=float, default=1.0) |
| parser.add_argument("--reconstruction-weight", type=float, default=0.0) |
| parser.add_argument("--reconstruction-size", type=int, default=32) |
| parser.add_argument("--selection-mode", choices=("raster_classification", "geometry"), default="raster_classification") |
| parser.add_argument( |
| "--raster-architecture", |
| choices=("fine_cross_attention_16x16_v6", "gated_fine_cross_attention_16x16_v7"), |
| ) |
| parser.add_argument("--classification-tolerance", type=float, default=0.005) |
| parser.add_argument("--holdout-tolerance", type=float, default=0.005) |
| parser.add_argument("--interpolation-alphas", default="0.125,0.25,0.5,1.0") |
| parser.add_argument("--seed", type=int, default=17) |
| parser.add_argument("--device", default="auto", help="auto|cpu|cuda[:index]") |
| args = parser.parse_args() |
| if min( |
| args.cycle_weight, args.preservation_weight, args.reconstruction_weight, |
| args.holdout_tolerance, args.classification_tolerance, |
| ) < 0: |
| raise ValueError("loss weight와 holdout tolerance는 0 이상이어야 합니다.") |
|
|
| random.seed(args.seed) |
| np.random.seed(args.seed) |
| torch.manual_seed(args.seed) |
| student_checkpoint = args.checkpoint |
| if args.raster_architecture: |
| |
| source_payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False) |
| initialized_model = MathInk06Model( |
| exact_classes=len(source_payload["exact_labels"]), family_classes=len(source_payload["family_labels"]), |
| hidden_size=int(source_payload.get("hidden_size", 128)), |
| hypotheses=int(source_payload.get("hypotheses", 4)), |
| raster_architecture=args.raster_architecture, |
| virtual_contract=str(source_payload.get("virtual_contract") or "legacy_v1"), |
| ) |
| initialized_model.load_state_dict(source_payload["state_dict"], strict=False) |
| source_payload["state_dict"] = initialized_model.state_dict() |
| source_payload["raster_architecture"] = args.raster_architecture |
| args.output.mkdir(parents=True, exist_ok=True) |
| student_checkpoint = args.output / "initialized_fine_checkpoint.pt" |
| torch.save(source_payload, student_checkpoint) |
| device = resolve_training_device06(args.device) |
| student = MathInk06Engine(student_checkpoint, device=device) |
| teacher = MathInk06Engine(args.checkpoint, device=str(student.device)) |
| print(json.dumps({"device": str(student.device), "cuda": torch.cuda.is_available()}), flush=True) |
| for parameter in student.model.parameters(): |
| parameter.requires_grad_(False) |
| trainable = [*student.model.raster_encoder.parameters(), *student.model.virtual_decoder.parameters()] |
| if student.model.auxiliary_virtual_decoder is not None: |
| trainable.extend(student.model.auxiliary_virtual_decoder.parameters()) |
| for parameter in trainable: |
| parameter.requires_grad_(True) |
| exact_to_index = {label: index for index, label in enumerate(student.labels)} |
| family_to_index = {label: index for index, label in enumerate(student.family_labels)} |
| sources = load_product_federation06( |
| registry_path=args.registry, commercial_path=args.commercial, hwrt_path=args.hwrt, |
| approval_path=args.approval, allowed_labels=student.labels, source_registry_path=args.source_registry, |
| ) |
| training, validation_groups, test_groups = _split_sources( |
| sources, seed=args.seed, train_max=args.max_train_per_source, |
| validation_max=args.max_validation_per_source, test_max=args.max_test_per_source, |
| ) |
| hard_labels = {value.strip() for value in args.hard_labels.split(",") if value.strip()} |
| unknown = hard_labels.difference(exact_to_index) |
| if unknown: |
| raise ValueError(f"378 vocabulary에 없는 hard label입니다: {sorted(unknown)}") |
| sampler = _hard_label_sampler( |
| training, seed=args.seed, hard_labels=hard_labels, multiplier=args.hard_label_multiplier, |
| ) |
| loader = DataLoader( |
| FederatedPairedInk06Dataset(training, exact_to_index, family_to_index), batch_size=args.batch_size, |
| sampler=sampler, num_workers=0, |
| ) |
| optimizer = torch.optim.AdamW(trainable, lr=args.learning_rate, weight_decay=1e-3) |
| baseline_validation = _evaluate_sources( |
| student, validation_groups, exact_to_index, family_to_index, args.batch_size, |
| ) |
| baseline_geometry = _evaluate_geometry_sources( |
| student, validation_groups, exact_to_index, family_to_index, args.batch_size, |
| ) |
| baseline_test = _evaluate_sources(student, test_groups, exact_to_index, family_to_index, args.batch_size) |
| best_score = ( |
| _macro_geometry_score(baseline_geometry) if args.selection_mode == "geometry" |
| else _macro_raster_score(baseline_validation) |
| ) |
| best_state = {key: value.detach().cpu().clone() for key, value in student.model.state_dict().items()} |
| anchor_state = {key: value.clone() for key, value in best_state.items()} |
| best_metrics = baseline_validation |
| best_geometry = baseline_geometry |
| best_epoch = 0 |
| best_alpha = 0.0 |
| alphas = tuple(float(value) for value in args.interpolation_alphas.split(",") if value.strip()) |
| if not alphas or any(not 0 < value <= 1 for value in alphas): |
| raise ValueError("interpolation alpha는 0보다 크고 1 이하여야 합니다.") |
| history: list[dict] = [] |
| for epoch in range(1, args.epochs + 1): |
| student.model.train() |
| totals: dict[str, float] = {} |
| seen = 0 |
| for online, raster, coordinates, states, target, family, _source in loader: |
| online, raster, coordinates, states, target, family = [ |
| value.to(student.device) for value in (online, raster, coordinates, states, target, family) |
| ] |
| optimizer.zero_grad(set_to_none=True) |
| with torch.inference_mode(): |
| teacher_output = teacher.model.forward_raster(raster) |
| loss, components = _losses( |
| student.model, online, raster, coordinates, states, target, family, |
| cycle_weight=args.cycle_weight, online_weight=0.0, multi_target=True, |
| teacher_output=teacher_output, preservation_weight=args.preservation_weight, |
| reconstruction_weight=args.reconstruction_weight, |
| reconstruction_size=args.reconstruction_size, |
| ) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(trainable, 1.0) |
| optimizer.step() |
| batch = len(target) |
| seen += batch |
| totals["loss"] = totals.get("loss", 0.0) + float(loss.detach()) * batch |
| for key, value in components.items(): |
| totals[key] = totals.get(key, 0.0) + value * batch |
| trained_state = {key: value.detach().cpu().clone() for key, value in student.model.state_dict().items()} |
| interpolation = [] |
| for alpha in alphas: |
| mixed = interpolate_state_dict06(anchor_state, trained_state, alpha=alpha) |
| student.model.load_state_dict(mixed) |
| metrics = _evaluate_sources(student, validation_groups, exact_to_index, family_to_index, args.batch_size) |
| geometry = _evaluate_geometry_sources( |
| student, validation_groups, exact_to_index, family_to_index, args.batch_size, |
| ) |
| score = ( |
| _macro_geometry_score(geometry) if args.selection_mode == "geometry" |
| else _macro_raster_score(metrics) |
| ) |
| tolerance = ( |
| args.classification_tolerance if args.selection_mode == "geometry" else args.holdout_tolerance |
| ) |
| guard = all( |
| metrics[source][metric] >= baseline_validation[source][metric] - tolerance |
| for source in metrics for metric in ("raster_top1", "raster_top5") |
| ) |
| interpolation.append({ |
| "alpha": alpha, "score": score, "holdout_guard": guard, |
| "validation": metrics, "geometry": geometry, |
| }) |
| if guard and score > best_score: |
| best_score, best_metrics, best_geometry, best_epoch, best_alpha = ( |
| score, metrics, geometry, epoch, alpha |
| ) |
| best_state = {key: value.clone() for key, value in mixed.items()} |
| student.model.load_state_dict(trained_state) |
| row = { |
| "epoch": epoch, "components": {key: value / max(seen, 1) for key, value in totals.items()}, |
| "interpolation": interpolation, |
| } |
| history.append(row) |
| print(json.dumps(row, ensure_ascii=False), flush=True) |
|
|
| student.model.load_state_dict(best_state) |
| final_test = _evaluate_sources(student, test_groups, exact_to_index, family_to_index, args.batch_size) |
| final_test_geometry = _evaluate_geometry_sources( |
| student, test_groups, exact_to_index, family_to_index, args.batch_size, |
| ) |
| payload = torch.load(student_checkpoint, map_location="cpu", weights_only=False) |
| payload["state_dict"] = best_state |
| payload["model_version"] = "aiflow-math-ink-0.6-federated-decoder1" |
| provenance = federation_provenance06(sources, args.source_registry) |
| payload.update(provenance) |
| payload["federated_decoder"] = { |
| "seed": args.seed, "hard_labels": sorted(hard_labels), "hard_label_multiplier": args.hard_label_multiplier, |
| "selected_epoch": best_epoch, "selected_interpolation_alpha": best_alpha, |
| "cycle_weight": args.cycle_weight, "preservation_weight": args.preservation_weight, |
| "reconstruction_weight": args.reconstruction_weight, "reconstruction_size": args.reconstruction_size, |
| "selection_mode": args.selection_mode, |
| } |
| payload["product_validation"] = False |
| args.output.mkdir(parents=True, exist_ok=True) |
| checkpoint = args.output / "math_ink_06_candidate.pt" |
| torch.save(payload, checkpoint) |
| report = { |
| "checkpoint": checkpoint.name, "bytes": checkpoint.stat().st_size, "seed": args.seed, |
| "device": str(student.device), |
| **provenance, |
| "train_samples": len(training), "source_count": len(sources), "hard_labels": sorted(hard_labels), |
| "baseline_validation": baseline_validation, "baseline_geometry": baseline_geometry, |
| "selected_validation": best_metrics, "selected_geometry": best_geometry, |
| "selected_epoch": best_epoch, "selected_interpolation_alpha": best_alpha, |
| "baseline_test": baseline_test, "test": final_test, "test_geometry": final_test_geometry, |
| "test_delta": { |
| source: {metric: final_test[source][metric] - baseline_test[source][metric] for metric in ( |
| "online_top1", "online_top5", "raster_top1", "raster_top5", |
| )} for source in final_test |
| }, |
| "history": history, "online_frozen": True, "product_validation": False, |
| } |
| (args.output / "federated_decoder_report.json").write_text( |
| json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", |
| ) |
| print(json.dumps({key: report[key] for key in report if key not in {"history"}}, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|