File size: 5,885 Bytes
cb634e7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""ModernBERT-based encoder for genomic MLM, augmented with Engram injection.

This wrapper adds Engram (conditional N-gram memory) deltas to specific
ModernBERT encoder layers via forward pre-hooks. It does NOT modify HF's
ModernBERT code, the ``engram`` module, or the ``model_modernbert`` module —
all three are reused unchanged.

Backbone parameter names match ``ModernBertWrapper`` exactly so a checkpoint
trained as ``model_type='modernbert'`` can be warm-started into this class
with ``strict=False``: only ``engram_blocks.*`` keys are missing.
"""
from __future__ import annotations

from dataclasses import replace
from typing import Sequence

import torch
import torch.nn as nn

from .engram import Engram, EngramConfig
from .model_modernbert import (
    ModernBertGenomicConfig,
    build_modernbert,
)


def _validate_layer_inject_ids(ids: Sequence[int], num_layers: int) -> list[int]:
    out: list[int] = []
    for idx, i in enumerate(ids):
        if isinstance(i, bool) or not isinstance(i, int):
            raise ValueError(
                f"layer_inject_ids[{idx}] = {i!r} is not an int (bools are rejected too)"
            )
        if i < 0 or i >= num_layers:
            raise ValueError(
                f"layer_inject_ids[{idx}] = {i} out of range [0, {num_layers})"
            )
        out.append(i)
    if len(set(out)) != len(out):
        raise ValueError(f"layer_inject_ids contains duplicates: {list(ids)}")
    return out


def _engram_key(i: int) -> str:
    return f"layer_{i:04d}"


class ModernBertEngramWrapper(nn.Module):
    """ModernBERT encoder for MLM with Engram modules injected at chosen layers.

    The wrapper:
      * Constructs a stock ``ModernBertForMaskedLM`` via :func:`build_modernbert`
        and stores it at ``self.model`` (identical attribute layout to
        :class:`ModernBertWrapper`, so state_dict keys are exactly the same
        for every backbone tensor).
      * Adds ``self.engram_blocks: nn.ModuleDict`` keyed by ``layer_{i:04d}``,
        one :class:`Engram` per index in ``engram_cfg.layer_inject_ids``.
      * Registers a ``forward_pre_hook(with_kwargs=True)`` on each target
        encoder layer that adds ``engram_blocks[layer_{i:04d}](hidden, input_ids)``
        to the layer's input ``hidden_states``.

    ``input_ids`` is stashed on ``self._current_input_ids`` at the start of
    each ``forward`` and never cleared inside the same call, so layers
    recomputed during backward (e.g. under ``torch.utils.checkpoint``) still
    see the right ids.
    """

    def __init__(self, cfg: ModernBertGenomicConfig, engram_cfg: EngramConfig):
        super().__init__()
        ids = _validate_layer_inject_ids(engram_cfg.layer_inject_ids, cfg.num_hidden_layers)
        engram_cfg = replace(
            engram_cfg,
            hidden_size=cfg.hidden_size,
            vocab_size=cfg.vocab_size,
            pad_id=cfg.pad_token_id,
        )
        self.cfg = cfg
        self.engram_cfg = engram_cfg
        # IDENTICAL outer attribute layout to ModernBertWrapper:
        self.model, self.hf_cfg = build_modernbert(cfg)
        # Sibling engram modules — NOT inside self.model, so backbone keys are unchanged.
        self.engram_blocks = nn.ModuleDict(
            {_engram_key(i): Engram(engram_cfg, layer_id=i) for i in ids}
        )
        for i in ids:
            self.model.model.layers[i].register_forward_pre_hook(
                self._make_hook(i), with_kwargs=True
            )
        self._current_input_ids: torch.Tensor | None = None

    @property
    def attn_implementation(self) -> str:
        impl = getattr(self.model.config, "_attn_implementation", None)
        if impl:
            return impl
        return getattr(self.model.config, "attn_implementation", "unknown")

    def num_params(self, only_trainable: bool = True) -> dict:
        backbone = 0
        memory = 0
        for name, p in self.named_parameters():
            if only_trainable and not p.requires_grad:
                continue
            n = p.numel()
            if "engram" in name and "embedding.emb" in name:
                memory += n
            else:
                backbone += n
        return {"backbone": backbone, "memory": memory, "total": backbone + memory}

    def _make_hook(self, layer_id: int):
        key = _engram_key(layer_id)
        # Cache to avoid one dict lookup per call.
        block_ref = self.engram_blocks[key]
        wrapper_self = self

        def hook(module, args, kwargs):
            ids = wrapper_self._current_input_ids
            if ids is None:
                return None  # cold-call safety (e.g. tracing); do nothing
            if args:
                hidden = args[0]
                delta = block_ref(hidden, input_ids=ids)
                new_args = (hidden + delta, *args[1:])
                return new_args, kwargs
            hidden = kwargs["hidden_states"]
            delta = block_ref(hidden, input_ids=ids)
            kwargs = {**kwargs, "hidden_states": hidden + delta}
            return args, kwargs

        return hook

    def forward(
        self,
        input_ids: torch.Tensor,
        labels: torch.Tensor | None = None,
        attention_mask: torch.Tensor | None = None,
    ) -> dict:
        # Overwrite-only on every forward; do NOT clear before/after model() so
        # that any backward-time recomputed forward (e.g. gradient checkpointing)
        # still reads the same ids that the original forward used.
        self._current_input_ids = input_ids
        if attention_mask is not None:
            attention_mask = attention_mask.to(torch.long)
        out = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            labels=labels,
            output_hidden_states=False,
        )
        return {"loss": out.loss, "logits": out.logits, "hidden": None}