Training Examples
Every workflow below has a complete runnable script in the examples directory.
Head Tuning
Task-head tuning documents the workflow, and tune_issue_head.py runs it end to end over a frozen backbone.
Multiclass Classification
Set num_labels and use problem_type="multiclass". Predictions from the named head
apply softmax and return one distribution per context:
head = model.fit_head(
labeled_issues,
task="issue_label",
num_labels=5,
problem_type="multiclass",
)
distribution = model.predict(batch, task_head="issue_label")
predicted_class = int(distribution.argmax())
Full-model Fine-tuning
Use RelationalTrainer to update the backbone and scalar decoder together. Gradient
accumulation supports effective batches larger than device memory:
args = RelationalTrainingArguments(
output_dir="models/churn",
num_train_epochs=3,
per_device_train_batch_size=8,
gradient_accumulation_steps=4, # optimizer sees 32 examples per step
)
RelationalTrainer(model=model, args=args, train_dataset=examples).train()
Script: finetune_churn.py
Evaluation During Training
Combine task metrics with ablation deltas to watch what fine-tuning changes:
from relational_transformers import (
AblationEvaluator,
BinaryClassificationEvaluator,
SequentialEvaluator,
)
evaluator = SequentialEvaluator([
BinaryClassificationEvaluator(validation_examples),
AblationEvaluator(validation_examples, ablations={"support": support_positions}),
])
metrics = evaluator(model)
Script: evaluate_churn.py
Custom PyTorch Loop
Call model.forward(batch, output="target_features") to train an arbitrary head, or
access model.model for complete control over optimization and mixed precision:
import torch
backbone = model.model
optimizer = torch.optim.AdamW(backbone.parameters(), lr=1e-5)
backbone.train()
for batch, labels in loader:
logits = backbone(batch.to(model.device), output="target_scores").scores
loss = loss_fn(logits, labels)
loss.backward()
optimizer.step()
optimizer.zero_grad(set_to_none=True)
The public loss modules from the Loss Overview plug into this loop unchanged.