Complete a PyTorch Training Pipeline to Extract Highlight Snippets from Text
Company: Figma
Role: Machine Learning Engineer
Category: Machine Learning
Difficulty: hard
Interview Round: Technical Screen
You are given a dataset for annotating a large number of text descriptions with highlights. Each sample is one text section and provides:
- the original text of the section;
- the highlighted text, the snippet of the section that should be highlighted;
- the start and end position indices of that snippet in the original text.
The data is split into a training set, a validation set and an evaluation set. Starting from a provided code skeleton in a notebook, complete a full PyTorch pipeline that learns to find the highlight in a text section: the model definition, the setup, the forward pass, training and evaluation. The skeleton has a stub for each step, but the implementation details are yours. The round lasts one hour, you can run and test code in the notebook, and you may look up PyTorch documentation.
The skeleton is not reproduced here, so define each piece yourself and state the interfaces you assume.
### Constraints and Clarifications
- Use PyTorch. Whether pretrained models or libraries beyond PyTorch are allowed is not stated; ask, and have a plan that works with PyTorch alone.
- Each sample provides one highlighted snippet and its positions.
### Clarifying Questions
- Are the start and end indices character offsets or token offsets, and is the end index inclusive or exclusive?
- At evaluation time, can a section contain more than one highlight, or none?
- May I use a pretrained text encoder and its tokenizer, or only PyTorch?
- Which metric decides success on the evaluation set: an exact match of the snippet, or partial credit for overlap?
- How long are the sections, and does the notebook have a GPU?
### Part 1 — Model definition
Define a model that takes a tokenized text section and produces the outputs you need to predict where the highlight is.
```hint Let the labels choose the head
Look at what the labels are before choosing an architecture: the output layer should let you compute a loss against them directly.
```
#### What This Part Should Cover
- A framing of the task that matches the labels
- An encoder that gives each token context from the rest of the section
- An output head with the right shape, and handling of padded positions
### Part 2 — Setup
Prepare everything training needs: turning raw samples into tensors and labels, batching, and the loss and optimizer.
```hint Check the offsets before you train
Verify on a few samples that slicing the original text with the given indices reproduces the highlighted text. That settles how the end index works.
```
```hint Characters versus tokens
The model sees tokens, but the labels refer to positions in the original text. You need a mapping in both directions.
```
#### What This Part Should Cover
- Tokenization that keeps each token's position in the original text
- Converting the snippet's positions into token labels, including truncation and spans that cannot be mapped
- Batching variable-length sections with padding and masks, and a vocabulary or tokenizer that does not peek at validation or evaluation data
- A loss and an optimizer suited to the model
### Part 3 — Forward pass
Implement the forward step used in training and in evaluation: from a batch to model outputs and, during training, a loss.
```hint Padding must lose
Padded positions should never be chosen as a prediction, and they should not contribute to the loss.
```
#### What This Part Should Cover
- Correct tensor shapes from input IDs to per-token scores
- Masking padded positions before the loss and before decoding
- A loss that covers both ends of the snippet, and a decoding step that always returns a valid snippet
### Part 4 — Training
Write the training loop.
```hint What the validation set is for
Decide how the validation set is used during training, and which version of the model you keep at the end.
```
#### What This Part Should Cover
- A correct loop: training mode, zeroing gradients, the backward pass, the optimizer step and device placement
- Stabilizers such as gradient clipping and a learning rate suited to the encoder
- Validation every epoch, keeping the best checkpoint and stopping early
### Part 5 — Evaluation
Evaluate the trained model on the evaluation set and report how well it finds the highlights.
```hint Score text, not indices
Compare the prediction with the label as text from the original section, so the metric does not depend on your tokenizer.
```
#### What This Part Should Cover
- Inference mode with gradients disabled
- Converting a decoded prediction back to a snippet of the original text
- Metrics with and without partial credit, computed over every evaluation sample
- Error analysis beyond a single number
### What a Strong Answer Covers
- A task framing that uses the start and end labels directly
- Alignment between text positions and tokens that is verified rather than assumed
- Masking and validity constraints applied consistently in the loss and in decoding
- A clean separation of the training, validation and evaluation data
- Runnable code that reaches an end-to-end result within the hour, completed in priority order
### Follow-up Questions
- How would you change the model if a section could contain several highlights, or none?
- Sections are longer than your encoder's maximum length. How do you train and predict?
- Your exact-match score is low but your overlap score is high. What does that tell you, and what would you try?
- How would you swap in a pretrained transformer encoder, and what changes in tokenization and training?
Overview: An ML coding question that asks you to complete a PyTorch pipeline for finding highlight snippets in text sections, given training, validation and evaluation sets with each snippet's start and end positions. It tests the model definition, data setup, the forward pass, the training loop and evaluation under time pressure.
Read the full Figma Machine Learning Engineer interview experience this question came from