Speech translation (seq2seq, variable length)¶
Variable-length "audio" features in one toy language, text tokens out in another. An encoder-decoder transformer, 173k params. CPU, no downloads, ~36 s.
gatle-ignite train --config=examples/translation/configs/translation_v0.py
python -m examples.translation.scripts.checks # the proof, run after training
gatle-ignite eval --config=examples/translation/configs/translation_v0.py --ckpt best
Nothing beyond the walk: dataset · model · prep_batch · optimizer · checkpoints · logging
| step | what this example changes |
|---|---|
| loss | dict_of_loss_params: ce (masked_ce) |
| metrics | val_metrics: exact (edit_distance), wer (edit_distance)score_name = valid/werscore_factor = -1 |
| running | grad_clip_norm = 1.0 |
The task is not a copy¶
Each of 12 source words has a prototype vector, uttered as 2-4 noisy frames. An utterance is 2-4 words, so sources are 4-16 frames and targets 4-6 tokens: variable on both sides. The reference translation is deterministic but reversed and permuted:
The reversal forces real cross-attention. A monotonic frame-i-to-token-i aligner cannot score on this, which is what makes a good result meaningful.
What this example demonstrates¶
- A custom
collate_fn: pass it throughds_paramsandbuild_dataloaderforwards it. It survives both DDP sharding paths, so a custom collate and exact eval sharding are orthogonal. - Masked loss:
ignore_index=PAD, so padding never trains toward<pad>. - Train/eval asymmetry: teacher forcing while training, greedy autoregressive decode at eval.
This needed zero step overrides:
prep_batchreadssplit, and the model branches on whether it was given a decoder input. - Lower-is-better scoring: WER selects the best checkpoint via
score_factor = -1.
from gatle_ignite import BaseTrainer, to_device
class Trainer(BaseTrainer):
def prep_batch(self, batch, split="train", **kwargs):
model_input = {
"src_feats": batch["src_feats"],
"src_pad_mask": batch["src_pad_mask"],
}
if split == "train":
model_input["tgt_in"] = batch["tgt_in"]
# Otherwise omit tgt_in, so the model's default (None) selects autoregressive decoding.
return to_device({"model_input": model_input, "targets": {"tgt_out": batch["tgt_out"]}})
The evidence¶
A falling teacher-forced loss proves nothing: a model can teacher-force well and decode into garbage. So the metrics score the decoded output.
- Positive control: the untrained model scores
wer 1.8381(above 1, because it emits garbage insertions) andexact 0.0000. - Padding really is masked: perturbing the logits by
25*N(0,1)at pad positions only moves the loss by exactly0.00e+00. The control that makes this non-vacuous: the same perturbation on an unmasked loss moves it2.938 -> 11.543. Gradients at PAD are0.000e+00; at real positions7.904e-03. - The decode is not secretly teacher-forced: rolling the source within the batch collapses exact
match
0.8984 -> 0.0078.
Two ways a check like these can lie
To test masking, perturb the pad logits, not the pad labels: changing a label stops that
position being ignored and renormalises reduction="mean". And to find the best checkpoint by
filename, take the largest score, not the smallest: the number in the filename is
score_factor * metric, so with score_factor = -1 the best is the least negative value.