Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user