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

This commit is contained in:
2026-07-30 05:04:34 -07:00
parent 9c43be6df1
commit 93e9c838bb
17 changed files with 1802 additions and 83 deletions
+132 -12
View File
@@ -22,12 +22,95 @@ DEFAULT_MODEL = "sentence-transformers/all-MiniLM-L6-v2"
# Reproducibility requires a model commit, not a mutable `main` branch.
DEFAULT_MODEL_REVISION = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41"
MAX_LENGTH = 128
HEAD_TAIL_SPECIAL_TOKENS = 3
HEAD_TOKENS = (MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS + 1) // 2
TAIL_TOKENS = MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS - HEAD_TOKENS
def prepare_text(prompt: str) -> str:
return normalize_prompt(prompt)
def encode_fixed_shape(
tokenizer: Any,
texts: Sequence[str],
torch: Any,
) -> dict[str, Any]:
"""Tokenize to 1x128 while retaining both context and a tail-buried request.
Pasted logs and stack traces frequently put the actual ask after the context. Plain
right truncation made generated boundary examples identical even when their final
request — and therefore their label — differed. Long inputs use BERT's sentence-pair
framing: [CLS] first 63 content tokens [SEP] last 62 content tokens [SEP].
"""
normalized = [prepare_text(text) for text in texts]
raw = tokenizer(
normalized,
add_special_tokens=False,
padding=False,
truncation=False,
return_attention_mask=False,
return_token_type_ids=False,
verbose=False,
)
if not isinstance(raw.get("input_ids"), list):
raise DataError("tokenizer did not return input_ids")
if tokenizer.pad_token_id is None:
raise DataError("purpose-lite tokenizer must define a padding token")
if tokenizer.cls_token_id is None or tokenizer.sep_token_id is None:
raise DataError("purpose-lite tokenizer must define BERT CLS and SEP tokens")
if tokenizer.padding_side != "right":
raise DataError("purpose-lite tokenizer must use right padding")
input_rows: list[list[int]] = []
mask_rows: list[list[int]] = []
type_rows: list[list[int]] = []
include_token_types = "token_type_ids" in tokenizer.model_input_names
single_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=False)
pair_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=True)
if pair_budget != HEAD_TOKENS + TAIL_TOKENS:
raise DataError(
"purpose-lite tokenizer special-token layout changed; expected three "
"tokens for head-tail inputs"
)
for content in raw["input_ids"]:
if len(content) <= single_budget:
first = content
second = None
else:
first = content[:HEAD_TOKENS]
second = content[-TAIL_TOKENS:]
if second is None:
input_ids = [tokenizer.cls_token_id] + first + [tokenizer.sep_token_id]
token_types = [0] * len(input_ids)
else:
input_ids = (
[tokenizer.cls_token_id]
+ first
+ [tokenizer.sep_token_id]
+ second
+ [tokenizer.sep_token_id]
)
token_types = [0] * (len(first) + 2) + [1] * (len(second) + 1)
if len(input_ids) > MAX_LENGTH:
raise DataError("fixed-shape tokenizer exceeded its 128-token contract")
padding = MAX_LENGTH - len(input_ids)
input_rows.append(input_ids + [tokenizer.pad_token_id] * padding)
mask_rows.append([1] * len(input_ids) + [0] * padding)
if include_token_types:
type_rows.append(token_types + [0] * padding)
encoded = {
"input_ids": torch.tensor(input_rows, dtype=torch.long),
"attention_mask": torch.tensor(mask_rows, dtype=torch.long),
}
if include_token_types:
encoded["token_type_ids"] = torch.tensor(type_rows, dtype=torch.long)
return encoded
def classification_metrics(
actual: Sequence[int], predicted: Sequence[int]
) -> dict[str, Any]:
@@ -267,6 +350,11 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
config.purpose_classifier_version = "purpose-lite-v1"
config.purpose_classifier_max_length = MAX_LENGTH
config.purpose_classifier_fixed_shape = [1, MAX_LENGTH]
config.purpose_classifier_truncation = {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
}
model.to(device)
class PromptDataset(Dataset):
@@ -282,13 +370,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]:
texts, labels = zip(*items)
encoded = tokenizer(
list(texts),
padding="max_length",
truncation=True,
max_length=MAX_LENGTH,
return_tensors="pt",
)
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
return encoded
@@ -311,6 +393,10 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
num_workers=args.workers,
pin_memory=device.type == "cuda",
)
validation_scorable = torch.tensor(
[record.get("slice") != "vague-eval" for record in validation_records],
dtype=torch.bool,
)
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
@@ -350,7 +436,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
logits, labels = _evaluate(torch, model, validation_loader, device)
predictions = logits.argmax(dim=-1).tolist()
metrics = classification_metrics(labels.tolist(), predictions)
scored_labels = labels[validation_scorable].tolist()
scored_predictions = [
prediction
for prediction, scorable in zip(
predictions, validation_scorable.tolist()
)
if scorable
]
metrics = classification_metrics(scored_labels, scored_predictions)
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = running_loss / len(train_loader)
history.append(metrics)
@@ -367,15 +461,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
logits, labels = _evaluate(torch, model, validation_loader, device)
temperature = _fit_temperature(torch, logits, labels)
temperature = _fit_temperature(
torch,
logits[validation_scorable],
labels[validation_scorable],
)
calibrated = torch.softmax(logits / temperature, dim=-1)
top = torch.topk(calibrated, k=2, dim=-1)
top_probabilities = top.values[:, 0].tolist()
margins = (top.values[:, 0] - top.values[:, 1]).tolist()
predictions = top.indices[:, 0].tolist()
correct = [
prediction == actual
for prediction, actual in zip(predictions, labels.tolist())
prediction == actual and scorable
for prediction, actual, scorable in zip(
predictions,
labels.tolist(),
validation_scorable.tolist(),
)
]
thresholds = choose_confidence_thresholds(
top_probabilities,
@@ -397,12 +499,30 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"baseModel": args.model,
"baseModelRevision": args.model_revision,
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"device": str(device),
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum().item()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum().item()
),
"bestValidationAccuracy": best_accuracy,
"bestValidation": classification_metrics(labels.tolist(), predictions),
"bestValidation": classification_metrics(
labels[validation_scorable].tolist(),
[
prediction
for prediction, scorable in zip(
predictions, validation_scorable.tolist()
)
if scorable
],
),
"history": history,
"calibration": calibration,
}