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