Project: variable-length text classifier¶
Train a small classifier from original image-description phrases. Follow vocabulary → token IDs → offsets → mean embeddings → class scores.
What to remember¶
Different-length sentences do not require a fixed padded matrix. EmbeddingBag accepts concatenated IDs plus the start offset of each sentence and returns one pooled vector per sentence.
This is an inexpensive baseline, not a contextual language model. It combines the course's tokenization, vocabulary, collation, pooling and class-weight ideas without external course helpers or pretrained downloads.
Important code¶
ids = torch.tensor([2, 3, 4, 5, 6], dtype=torch.long)
offsets = torch.tensor([0, 2], dtype=torch.long)
embedding = nn.EmbeddingBag(vocab_size, 12, mode="mean", padding_idx=0)
head = nn.Linear(12, 2)
logits = head(embedding(ids, offsets)) # [2, 2]
The first sentence is [2,3], the second is [4,5,6]. Their lengths differ, but both produce a 12-value vector. Offsets use start positions; the default API does not require a final ending offset.
lengths 2 and 3five token IDsoffsets [0,2]two starts, no padding[2,12]one mean vector per sentence[2,2]two class scoresscalarclass-weighted cross-entropyIllustration: shapes and operations, not measured model performance.
| Choice | This example | Why it matters |
|---|---|---|
| Vocabulary | Sorted words from training text only | Prevents validation text from defining the representation. |
| Reserved IDs | Padding 0, unknown 1 |
Unknown words and padding mean different things. |
| Empty input | One unknown token | Makes the empty-input policy explicit. |
| Collation | Flat IDs, start offsets and integer labels | Keeps variable-length batching on CPU. |
| Class weights | Inverse training counts, aligned with IDs | Gives the rarer class more weight in the loss. |
| Pooling | Mean | Balances length but discards word order. |
The tiny training set has six positive and two negative descriptions. It computes weights [2.0, 2/3] in label order. More weight changes the loss priorities; it does not guarantee better minority-class predictions.
Check the representation¶
If you reverse the words in a sentence, its mean-pooled vector stays the same. “Dog bites person” and “person bites dog” therefore cannot be distinguished by this baseline.
A padded nn.Embedding alternative needs a mask when computing a manual mean. Compare its padding cost below; the script itself uses offsets and no padding:
Run and inspect¶
python examples/text_bags.py
python examples/text_bags.py --epochs 40 --output artifacts/my-text-run
The offline CPU run saves predictions.json and text-model.pt under the output directory. Inspect the token map, class weights, validation history and predictions for four new combinations of words. The checkpoint includes the vocabulary, label names and tokenization/empty-input policy.
These toy phrases demonstrate the data flow. Their validation scores are not a sentiment benchmark and do not prove useful performance on natural language. For a real application, enlarge the dataset and inspect per-class recall and macro F1.
Try changing: add a new training phrase, rerun vocabulary construction, then compare a new word with an unknown word. If order changes the correct label, move to an order-aware or contextual model rather than tuning a bag of words indefinitely.
Review tokenization and embeddings · Review classifiers and fine-tuning
Inspect a retained tiny execution¶
Measured toy execution, CPU, 20 epochs, seed 53. Eight training phrases and four validation phrases are too small to establish general sentiment quality. Here the labels describe image quality.
| Word | ID | First four values of its 12-value learned row |
|---|---|---|
<pad> |
0 | 0.000, 0.000, 0.000, 0.000 |
<unk> |
1 | -0.959, 0.895, -0.575, 0.478 |
bad |
2 | 0.441, -0.121, 0.357, -1.113 |
blurry |
3 | -0.764, -1.040, 0.530, 1.699 |
bright |
4 | -0.896, 2.745, -1.294, 0.934 |
clear |
5 | 2.665, 1.200, -0.462, 0.496 |
dark |
6 | -2.493, 0.230, 1.268, 1.275 |
good |
7 | 0.557, 0.951, 1.377, 0.468 |
image |
8 | -0.549, 0.536, 0.037, 0.650 |
photo |
9 | -1.242, 0.326, -0.902, 0.074 |
picture |
10 | 0.071, 0.787, -1.774, -0.047 |
sharp |
11 | 1.112, 0.490, -0.195, 0.725 |
The first two phrases clear image and sharp picture become flat IDs [5, 8, 11, 10] and offsets [0, 2]. Each offset starts a segment; the final segment runs to the end. Mean pooling turns each segment into one 12-value row; the classifier returns two logits per phrase. No PAD rows are needed for this bag representation.
Training class counts are [2, 6]; inverse-frequency weights are [2.0, 0.6667]. For class-index cross-entropy, the mean divides weighted losses by the sum of target weights, not by batch size. PyTorch loss contract.
| Validation phrase | True label | Actual prediction |
|---|---|---|
| good sharp image | 1 | 1 |
| bad dark picture | 0 | 0 |
| clear bright photo | 1 | 1 |
| blurry image | 0 | 0 |
Labels: 0 = poor-image-description, 1 = good-image-description. Unknown input unseenword becomes [1] and predicts 1 in this run. The unknown row has no supervised training examples here; this prediction is not evidence of understanding an unknown word. Empty input also becomes <unk>.
clear image and image clear have maximum logit difference 0.0: a mean bag loses order even though the original texts differ. Their actual logits are [-2.6688294410705566, 2.805745840072632]. Use an order-aware model when the distinction changes meaning. Retained vocabulary, IDs, offsets, embeddings and predictions.
Complete source¶
Open the maintained runnable script
"""Original tiny text classifier: vocabulary, offsets, class weights and saved contract."""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
import torch
from torch import nn
from torch.utils.data import DataLoader
from recall_patterns import collate_bags
TRAIN = [
("clear image", 1), ("sharp picture", 1), ("good bright photo", 1),
("clear sharp photo", 1), ("good image", 1), ("bright picture", 1),
("blurry dark image", 0), ("bad blurry photo", 0),
]
VALIDATION = [("good sharp image", 1), ("bad dark picture", 0),
("clear bright photo", 1), ("blurry image", 0)]
def tokenize(text):
return re.findall(r"[a-z]+", text.lower())
def vocabulary(samples):
words = sorted({word for text, _ in samples for word in tokenize(text)})
return {"<pad>": 0, "<unk>": 1, **{word: i + 2 for i, word in enumerate(words)}}
def encode(text, vocab):
ids = [vocab.get(word, 1) for word in tokenize(text)]
return torch.tensor(ids or [1], dtype=torch.long)
class BagClassifier(nn.Module):
def __init__(self, vocab_size):
super().__init__()
self.embedding = nn.EmbeddingBag(vocab_size, 12, mode="mean", padding_idx=0)
self.head = nn.Linear(12, 2)
def forward(self, ids, offsets):
return self.head(self.embedding(ids, offsets))
def run(output: Path, epochs=20):
torch.set_num_threads(2)
torch.manual_seed(53)
vocab = vocabulary(TRAIN) # no validation words are used to build the vocabulary
encoded = [(encode(text, vocab), label) for text, label in TRAIN]
loader = DataLoader(encoded, batch_size=3, shuffle=True, collate_fn=collate_bags,
generator=torch.Generator().manual_seed(59))
counts = torch.bincount(torch.tensor([label for _, label in TRAIN]), minlength=2)
weights = counts.sum() / (2 * counts.float()) # [class 0, class 1]
model = BagClassifier(len(vocab))
optimizer = torch.optim.Adam(model.parameters(), lr=0.03)
loss_fn = nn.CrossEntropyLoss(weight=weights)
history = []
for epoch in range(epochs):
model.train()
for ids, offsets, labels in loader:
optimizer.zero_grad(set_to_none=True)
logits = model(ids, offsets)
loss = loss_fn(logits, labels)
loss.backward()
optimizer.step()
model.eval()
with torch.no_grad():
ids, offsets, labels = collate_bags(
[(encode(text, vocab), label) for text, label in VALIDATION])
predicted = model(ids, offsets).argmax(1)
accuracy = (predicted == labels).float().mean().item()
history.append({"epoch": epoch + 1, "validation_accuracy": accuracy})
predictions = [{"text": text, "target": label, "predicted": int(pred)}
for (text, label), pred in zip(VALIDATION, predicted)]
with torch.no_grad():
unknown_ids = encode("unseenword", vocab)
unknown_prediction = int(model(unknown_ids, torch.tensor([0])).argmax(1))
forward = model(encode("clear image", vocab), torch.tensor([0]))
reversed_words = model(encode("image clear", vocab), torch.tensor([0]))
example_ids, example_offsets, _ = collate_bags(encoded[:2])
output.mkdir(parents=True, exist_ok=True)
report = {
"data": "tiny original toy phrases; not a sentiment benchmark",
"vocabulary": vocab, "labels": ["poor-image-description", "good-image-description"],
"tokenization": "lowercase ASCII word regex", "empty_input": "<unk>",
"pooling": "mean; no word order; no padding required",
"class_weights": weights.tolist(), "history": history, "predictions": predictions,
"encoded_batch": {"texts": [text for text, _ in TRAIN[:2]], "ids": example_ids.tolist(), "offsets": example_offsets.tolist()},
"unknown_input": {"text": "unseenword", "ids": unknown_ids.tolist(), "predicted": unknown_prediction},
"order_check": {"texts": ["clear image", "image clear"], "logits": [forward.tolist()[0], reversed_words.tolist()[0]], "max_logit_difference": (forward - reversed_words).abs().max().item()},
"embedding_preview": {word: model.embedding.weight[index, :4].detach().tolist() for word, index in vocab.items()},
}
torch.save({"state_dict": model.state_dict(), "metadata": report}, output / "text-model.pt")
(output / "predictions.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
print(json.dumps(predictions, indent=2))
return report
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--output", type=Path, default=Path("artifacts/text-bags"))
parser.add_argument("--epochs", type=int, default=20)
args = parser.parse_args()
if args.epochs < 1:
parser.error("--epochs must be positive")
run(args.output, args.epochs)
Open the shared variable-length collator
def collate_bags(samples):
# Each sample is (1D token-id tensor, integer label); keep batching on CPU.
sequences, labels = zip(*samples)
sequences = [s if s.numel() else torch.tensor([1]) for s in sequences]
lengths = torch.tensor([s.numel() for s in sequences], dtype=torch.long)
offsets = torch.cat([torch.zeros(1, dtype=torch.long), lengths.cumsum(0)[:-1]])
return torch.cat(sequences), offsets, torch.tensor(labels, dtype=torch.long)
API reference: EmbeddingBag inputs, offsets and padding.