File size: 13,134 Bytes
a75ccfd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
#!/usr/bin/env python3
"""DistilBERT + a linear classification head, in either ``lora`` or ``full`` mode.

Both modes use the identical architecture. The only difference is which
parameters carry ``requires_grad=True`` β€” which is precisely the comparison this
project is about.

Run ``python -m models.classifier`` to print both modes side by side and see the
parameter counts without training anything.
"""

from __future__ import annotations

import json
import logging
from pathlib import Path
from typing import Any, Literal

import torch
import torch.nn as nn
from transformers import AutoConfig, AutoModel

try:
    from models.lora import inject_lora, mark_only_lora_as_trainable, parameter_report
except ImportError:  # pragma: no cover - fallback for direct script execution
    from lora import inject_lora, mark_only_lora_as_trainable, parameter_report  # type: ignore

log = logging.getLogger("classifier")

Mode = Literal["lora", "full"]

DEFAULT_MODEL = "distilbert-base-uncased"

#: DistilBERT names its attention projections q_lin / k_lin / v_lin / out_lin.
#: Query and value only β€” see models/lora.py for why.
LORA_TARGETS: tuple[str, ...] = ("q_lin", "v_lin")


class TextClassifier(nn.Module):
    """Classifies a short text into one of ``num_labels`` classes.

    Architecture::

        input_ids (B, L)
          -> DistilBERT encoder            -> (B, L, 768)
          -> take position 0, the [CLS]    -> (B, 768)
          -> dropout
          -> Linear(768, num_labels)       -> (B, num_labels)   logits

    A deliberate simplification: HuggingFace's own
    ``DistilBertForSequenceClassification`` inserts a 768x768 ``pre_classifier``
    layer before the final one. That adds ~590k trainable parameters to the head β€”
    which would be **four times larger than all the LoRA adapters combined** and
    would completely distort the headline "1% of parameters" comparison. A single
    linear head (59k params at 77 classes) keeps the measurement honest.
    """

    def __init__(
        self,
        model_name: str = DEFAULT_MODEL,
        num_labels: int = 77,
        mode: Mode = "lora",
        lora_r: int = 8,
        lora_alpha: int = 16,
        lora_dropout: float = 0.0,
        head_dropout: float = 0.1,
        class_weights: torch.Tensor | None = None,
        label_names: list[str] | None = None,
    ) -> None:
        """
        Args:
            model_name: Pretrained checkpoint to load.
            num_labels: Number of classes. 77 for banking77.
            mode: ``"lora"`` freezes the base and adapts query/value projections.
                ``"full"`` trains every parameter.
            lora_r: LoRA rank (ignored in full mode).
            lora_alpha: LoRA scaling numerator (ignored in full mode).
            lora_dropout: Dropout inside the LoRA path (ignored in full mode).
            head_dropout: Dropout before the classification head.
            class_weights: Optional ``(num_labels,)`` tensor for
                ``CrossEntropyLoss``. See ``data.dataset.class_weights``.
            label_names: Ordered class names. Stored in the checkpoint so the
                demo and evaluator resolve predictions to names without needing
                the dataset β€” the single most likely source of silently
                mislabelled output if left to be re-derived.
        """
        super().__init__()
        if mode not in ("lora", "full"):
            raise ValueError(f"mode must be 'lora' or 'full', got {mode!r}")

        self.model_name = model_name
        self.num_labels = num_labels
        self.mode = mode
        self.label_names = list(label_names) if label_names else None
        self.lora_r = lora_r
        self.lora_alpha = lora_alpha

        config = AutoConfig.from_pretrained(model_name)
        self.encoder = AutoModel.from_pretrained(model_name)
        hidden_size = config.dim if hasattr(config, "dim") else config.hidden_size

        self.dropout = nn.Dropout(head_dropout)
        self.classifier = nn.Linear(hidden_size, num_labels)

        self.n_adapted = 0
        if mode == "lora":
            self.n_adapted = inject_lora(
                self.encoder,
                target_names=LORA_TARGETS,
                r=lora_r,
                alpha=lora_alpha,
                dropout=lora_dropout,
            )
            log.info("Injected LoRA (r=%d) into %d projections", lora_r, self.n_adapted)
            mark_only_lora_as_trainable(self, also_train=("classifier",))

        # register_buffer, not a plain attribute: buffers move with .to(device)
        # but are not parameters, so the optimizer ignores them.
        if class_weights is not None:
            self.register_buffer("class_weights", class_weights)
        else:
            self.class_weights = None

    def forward(
        self,
        input_ids: torch.Tensor,
        attention_mask: torch.Tensor,
        labels: torch.Tensor | None = None,
    ) -> dict[str, torch.Tensor | None]:
        """Run a forward pass and, if labels are given, compute the loss.

        Args:
            input_ids: ``(B, L)`` int64 token ids.
            attention_mask: ``(B, L)`` int64, 1 for real tokens and 0 for padding.
            labels: ``(B,)`` int64 class ids in ``[0, num_labels)``, or None at
                inference time.

        Returns:
            ``{"logits": (B, num_labels), "loss": scalar or None}``.
        """
        # (B, L) -> (B, L, 768). The attention mask keeps padded positions from
        # contributing to any token's representation.
        hidden = self.encoder(
            input_ids=input_ids, attention_mask=attention_mask
        ).last_hidden_state

        # DistilBERT has no pooler of its own, so pool manually: position 0 is
        # the [CLS] token, whose representation attends over the whole sequence
        # and is the conventional sentence-level summary.
        # (B, L, 768) -> (B, 768)
        pooled = hidden[:, 0]

        logits = self.classifier(self.dropout(pooled))  # (B, 768) -> (B, num_labels)

        loss = None
        if labels is not None:
            # CrossEntropyLoss takes RAW logits, not softmax probabilities β€” it
            # applies log_softmax internally. Feeding it softmax output is a
            # classic silent bug: it still trains, just badly.
            loss_fn = nn.CrossEntropyLoss(weight=self.class_weights)
            loss = loss_fn(logits, labels)

        return {"logits": logits, "loss": loss}

    @torch.no_grad()
    def predict(
        self, input_ids: torch.Tensor, attention_mask: torch.Tensor
    ) -> tuple[torch.Tensor, torch.Tensor]:
        """Return ``(predicted_ids (B,), probabilities (B, num_labels))``.

        ``@torch.no_grad()`` disables the autograd graph β€” at inference it is
        pure overhead in both memory and time.
        """
        self.eval()
        logits = self.forward(input_ids, attention_mask)["logits"]
        probs = torch.softmax(logits, dim=-1)
        return probs.argmax(dim=-1), probs

    def trainable_parameter_report(self) -> dict[str, Any]:
        """The project's key measurement: how much of the model is actually trained."""
        report = parameter_report(self)
        report.update(
            {
                "mode": self.mode,
                "lora_rank": self.lora_r if self.mode == "lora" else None,
                "adapted_layers": self.n_adapted,
            }
        )
        return report

    # ------------------------------------------------------------------
    # Checkpointing
    # ------------------------------------------------------------------

    def adapter_state_dict(self) -> dict[str, torch.Tensor]:
        """Only the tensors that actually changed during LoRA training.

        The frozen encoder is byte-identical to the public checkpoint, so saving
        it would be storing a copy of something already on the Hub. Keeping just
        the adapters and head takes the artifact from ~265MB to ~200KB, which is
        what makes the free Hugging Face Space deploy practical (see DEPLOY.md).
        """
        return {
            name: param.detach().cpu().clone()
            for name, param in self.state_dict().items()
            if "lora_A" in name or "lora_B" in name or name.startswith("classifier")
        }

    def save(self, path: str | Path) -> None:
        """Save a checkpoint: adapters only in LoRA mode, everything in full mode."""
        path = Path(path)
        path.parent.mkdir(parents=True, exist_ok=True)

        state = (
            self.adapter_state_dict()
            if self.mode == "lora"
            else {k: v.detach().cpu() for k, v in self.state_dict().items()}
        )
        torch.save(
            {
                "state_dict": state,
                "config": {
                    "model_name": self.model_name,
                    "num_labels": self.num_labels,
                    "mode": self.mode,
                    "lora_r": self.lora_r,
                    "lora_alpha": self.lora_alpha,
                    "label_names": self.label_names,
                },
            },
            path,
        )
        size_kb = path.stat().st_size / 1024
        log.info("Saved %s checkpoint to %s (%.0f KB)", self.mode, path, size_kb)

    @classmethod
    def load(cls, path: str | Path, device: torch.device | str = "cpu") -> "TextClassifier":
        """Rebuild a trained model from a checkpoint.

        The saved config records the mode and rank, so the architecture is
        reconstructed correctly before the weights are loaded in. Getting the
        rank wrong would produce a shape mismatch here rather than silently
        wrong predictions β€” which is the point of storing it.
        """
        # weights_only=False: the checkpoint holds a config dict alongside the
        # tensors. Only load checkpoints you produced yourself.
        payload = torch.load(path, map_location=device, weights_only=False)
        config = payload["config"]

        model = cls(
            model_name=config["model_name"],
            num_labels=config["num_labels"],
            mode=config["mode"],
            lora_r=config["lora_r"],
            lora_alpha=config["lora_alpha"],
            label_names=config.get("label_names"),
        )
        # strict=False for LoRA: the checkpoint intentionally omits the frozen
        # encoder, which came from the pretrained download instead.
        missing, unexpected = model.load_state_dict(
            payload["state_dict"], strict=(config["mode"] == "full")
        )
        if unexpected:
            raise ValueError(f"Checkpoint has unexpected keys: {unexpected[:5]}")
        if config["mode"] == "lora":
            mark_only_lora_as_trainable(model, also_train=("classifier",))

        return model.to(device)


def _demo() -> None:
    """Print both modes side by side. Run: ``python -m models.classifier``."""
    logging.basicConfig(level=logging.INFO, format="%(levelname)-7s %(message)s")
    for noisy in ("httpx", "urllib3", "filelock", "huggingface_hub"):
        logging.getLogger(noisy).setLevel(logging.WARNING)

    print("\n" + "=" * 68)
    print("  PARAMETER COMPARISON β€” no training, just architecture")
    print("=" * 68)

    try:
        from data.dataset import load_labels
        labels = load_labels()
    except Exception:  # noqa: BLE001 - demo works even before download.py runs
        labels = [f"class_{i}" for i in range(77)]
    print(f"  (using {len(labels)} classes)")

    rows = []
    for mode in ("full", "lora"):
        model = TextClassifier(mode=mode, num_labels=len(labels), label_names=labels)
        report = model.trainable_parameter_report()
        rows.append(report)
        label = f"{mode}" + (f" (r={report['lora_rank']})" if mode == "lora" else "")
        print(f"\n  {label}")
        print("  " + "-" * 56)
        print(f"    trainable : {report['trainable_params']:>12,}")
        print(f"    frozen    : {report['frozen_params']:>12,}")
        print(f"    total     : {report['total_params']:>12,}")
        print(f"    trainable : {report['trainable_pct']:>11.3f}% of total")
        if mode == "lora":
            print(f"    adapted   : {report['adapted_layers']} projections "
                  f"(query + value across 6 layers)")

    ratio = rows[0]["trainable_params"] / rows[1]["trainable_params"]
    print("\n  " + "-" * 56)
    print(f"  Full fine-tuning trains {ratio:,.0f}x more parameters than LoRA.")

    # Verify the zero-init claim rather than asserting it in a comment.
    lora_model = TextClassifier(mode="lora", num_labels=len(labels))
    first_adapter = next(
        m for m in lora_model.encoder.modules() if type(m).__name__ == "LoRALinear"
    )
    print(f"\n  Sanity check at initialization:")
    print(f"    lora_B all zeros : {bool((first_adapter.lora_B == 0).all())}  "
          f"(so B@A = 0, model == pretrained)")
    print(f"    lora_A all zeros : {bool((first_adapter.lora_A == 0).all())}  "
          f"(must be False, or no gradient could flow)")
    print("\n" + "=" * 68 + "\n")


if __name__ == "__main__":
    _demo()