Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 17:28:44 -07:00
parent 93e9c838bb
commit e4403b570d
7 changed files with 408 additions and 19 deletions
+103 -13
View File
@@ -31,6 +31,12 @@ def prepare_text(prompt: str) -> str:
return normalize_prompt(prompt)
def training_weight(record: dict[str, Any], boundary_weight: float) -> float:
"""Return the loss weight for one training record."""
return boundary_weight if record.get("slice") == "boundary" else 1.0
def encode_fixed_shape(
tokenizer: Any,
texts: Sequence[str],
@@ -284,6 +290,7 @@ def _evaluate(torch: Any, model: Any, loader: Any, device: Any) -> tuple[Any, An
with torch.inference_mode():
for batch in loader:
labels = batch.pop("labels")
batch.pop("sample_weights", None)
inputs = {key: value.to(device) for key, value in batch.items()}
logits = model(**inputs).logits.cpu()
all_logits.append(logits)
@@ -317,6 +324,17 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
validation_records = validation_records[: args.max_validation_records]
output_dir: Path = args.output_dir
local_model = Path(args.model).expanduser()
if local_model.exists():
try:
local_model.resolve().relative_to(output_dir.resolve())
except ValueError:
pass
else:
raise DataError(
"local --model must not be inside --output-dir; overwrite could "
"destroy the continuation checkpoint"
)
if output_dir.exists() and any(output_dir.iterdir()):
if not args.overwrite_output:
raise DataError(
@@ -329,16 +347,22 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
device = _select_device(torch, args.device)
label_to_id = {label: index for index, label in enumerate(LABELS)}
id_to_label = {index: label for label, index in label_to_id.items()}
model_revision = None if local_model.exists() else args.model_revision
pretrained_options = (
{"local_files_only": True}
if local_model.exists()
else {"revision": model_revision}
)
tokenizer = AutoTokenizer.from_pretrained(
args.model, revision=args.model_revision, use_fast=True
args.model, use_fast=True, **pretrained_options
)
model = AutoModelForSequenceClassification.from_pretrained(
args.model,
revision=args.model_revision,
num_labels=len(LABELS),
label2id=label_to_id,
id2label=id_to_label,
ignore_mismatched_sizes=True,
**pretrained_options,
)
config = model.config
if getattr(config, "hidden_size", None) != 384 or getattr(
@@ -364,14 +388,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
def __len__(self) -> int:
return len(self.records)
def __getitem__(self, index: int) -> tuple[str, int]:
def __getitem__(self, index: int) -> tuple[str, int, float]:
record = self.records[index]
return prepare_text(record["prompt"]), label_to_id[record["purpose"]]
return (
prepare_text(record["prompt"]),
label_to_id[record["purpose"]],
training_weight(record, args.boundary_weight),
)
def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]:
texts, labels = zip(*items)
def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]:
texts, labels, weights = zip(*items)
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float)
return encoded
generator = torch.Generator()
@@ -397,6 +426,36 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
[record.get("slice") != "vague-eval" for record in validation_records],
dtype=torch.bool,
)
initial_logits, initial_labels = _evaluate(
torch,
model,
validation_loader,
device,
)
initial_predictions = initial_logits.argmax(dim=-1).tolist()
initial_metrics = classification_metrics(
initial_labels[validation_scorable].tolist(),
[
prediction
for prediction, scorable in zip(
initial_predictions,
validation_scorable.tolist(),
)
if scorable
],
)
best_accuracy = initial_metrics["accuracy"]
epochs_without_improvement = 0
stopped_early = False
history = []
best_dir = output_dir / "model"
model.save_pretrained(best_dir, safe_serialization=True)
tokenizer.save_pretrained(best_dir)
print(
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
f"macro_recall={initial_metrics['macroRecall']:.4%}",
flush=True,
)
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
@@ -411,17 +470,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
num_training_steps=total_steps,
)
best_accuracy = -1.0
history = []
best_dir = output_dir / "model"
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
model.train()
optimizer.zero_grad(set_to_none=True)
running_loss = 0.0
for step, batch in enumerate(train_loader, 1):
batch = {key: value.to(device) for key, value in batch.items()}
loss = model(**batch).loss / args.gradient_accumulation_steps
labels = batch.pop("labels").to(device)
sample_weights = batch.pop("sample_weights").to(device)
inputs = {key: value.to(device) for key, value in batch.items()}
per_record_loss = torch.nn.functional.cross_entropy(
model(**inputs).logits,
labels,
reduction="none",
)
loss = (
(per_record_loss * sample_weights).sum() / sample_weights.sum()
) / args.gradient_accumulation_steps
loss.backward()
running_loss += float(loss.item()) * args.gradient_accumulation_steps
should_update = (
@@ -454,10 +519,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
f"macro_recall={metrics['macroRecall']:.4%}",
flush=True,
)
if metrics["accuracy"] > best_accuracy:
improvement = metrics["accuracy"] - best_accuracy
if improvement > args.minimum_improvement:
best_accuracy = metrics["accuracy"]
epochs_without_improvement = 0
model.save_pretrained(best_dir, safe_serialization=True)
tokenizer.save_pretrained(best_dir)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no validation improvement "
f"greater than {args.minimum_improvement:.4%} for "
f"{args.early_stopping_patience} epoch(s)",
flush=True,
)
break
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
logits, labels = _evaluate(torch, model, validation_loader, device)
@@ -497,7 +575,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
metrics = {
"modelVersion": "purpose-lite-v1",
"baseModel": args.model,
"baseModelRevision": args.model_revision,
"baseModelRevision": model_revision or "local-checkpoint",
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
@@ -507,12 +585,16 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"device": str(device),
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"boundaryTrainingWeight": args.boundary_weight,
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum().item()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum().item()
),
"bestValidationAccuracy": best_accuracy,
"initialValidation": initial_metrics,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"bestValidation": classification_metrics(
labels[validation_scorable].tolist(),
[
@@ -555,6 +637,9 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--warmup-ratio", type=float, default=0.1)
parser.add_argument("--max-grad-norm", type=float, default=1.0)
parser.add_argument("--workers", type=int, default=0)
parser.add_argument("--early-stopping-patience", type=int, default=2)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
parser.add_argument("--boundary-weight", type=float, default=1.0)
parser.add_argument("--high-precision", type=float, default=0.98)
parser.add_argument("--accepted-precision", type=float, default=0.95)
parser.add_argument("--max-train-records", type=int)
@@ -576,10 +661,15 @@ def main(argv: Sequence[str] | None = None) -> int:
"batch_size",
"eval_batch_size",
"gradient_accumulation_steps",
"early_stopping_patience",
):
_positive(parser, f"--{name.replace('_', '-')}", getattr(args, name))
if not 0.0 <= args.warmup_ratio < 1.0:
parser.error("--warmup-ratio must be in [0, 1)")
if args.minimum_improvement < 0.0:
parser.error("--minimum-improvement must be non-negative")
if args.boundary_weight <= 0.0:
parser.error("--boundary-weight must be positive")
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
parser.error(
"precision targets must satisfy 0 < accepted <= high <= 1"