aiflow-math-ink-06-intermediate / scripts /train_math_ink_06_federated_decoder.py
cwLeeDev's picture
Publish clean federation GPU ablation and bottleneck audit
0eef691 verified
Raw
History Blame Contribute Delete
18.8 kB
"""승인 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:
# 기존 online/head를 그대로 두고 새 fine raster branch의 추가 weight만 초기화한다.
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()