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.