File size: 8,587 Bytes
76bab86
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Residual MultiConvAdapter on the frozen Cohere Conformer.

Kernels K={7,15,23,31} + concat_fusion from MULTI-CONVFORMER
(Prabhu et al. 2024). Skip the bottom third of layers. Zero-init the
up-projection so the first step is identity.
"""

from __future__ import annotations

from typing import Iterable

import torch
import torch.nn as nn

DEFAULT_KERNELS = (7, 15, 23, 31)
DEFAULT_FUSION = "concat_fusion"
DEFAULT_MERGE_KERNEL = 31


class MultiConvAdapter(nn.Module):
    def __init__(
        self,
        d_model: int,
        bottleneck: int = 64,
        kernels: Iterable[int] = DEFAULT_KERNELS,
        dropout: float = 0.1,
        fusion: str = DEFAULT_FUSION,
        merge_kernel: int = DEFAULT_MERGE_KERNEL,
    ):
        super().__init__()
        kernels = tuple(int(k) for k in kernels)
        if not kernels:
            raise ValueError("Need at least one convolution kernel")
        if any(k < 1 or k % 2 == 0 for k in kernels):
            raise ValueError(f"Kernels must be odd and positive, got {kernels}")
        if bottleneck < len(kernels) or bottleneck % 2 != 0:
            raise ValueError(f"bottleneck must be even and >= n_kernels, got {bottleneck}")
        if fusion not in {"sum", "weighted_sum", "concat", "concat_fusion"}:
            raise ValueError(f"Unknown fusion={fusion}")
        if fusion in {"concat", "concat_fusion"} and bottleneck % len(kernels) != 0:
            raise ValueError(
                f"concat fusion needs bottleneck ({bottleneck}) divisible by "
                f"{len(kernels)} kernels"
            )

        self.kernels = kernels
        self.fusion = fusion
        self.merge_kernel = int(merge_kernel)
        self.norm = nn.LayerNorm(d_model)
        self.down = nn.Linear(d_model, 2 * bottleneck)
        self.act = nn.GELU()
        self.gate_norm = nn.LayerNorm(bottleneck)
        n_kernels = len(kernels)

        if fusion in {"sum", "weighted_sum"}:
            self.convs = nn.ModuleList(
                [
                    nn.Conv1d(
                        bottleneck,
                        bottleneck,
                        kernel_size=k,
                        padding=(k - 1) // 2,
                        groups=bottleneck,
                    )
                    for k in kernels
                ]
            )
        else:
            per = bottleneck // n_kernels
            self.convs = nn.ModuleList(
                [
                    nn.Conv1d(
                        bottleneck,
                        per,
                        kernel_size=k,
                        padding=(k - 1) // 2,
                        groups=per,
                    )
                    for k in kernels
                ]
            )

        if fusion == "weighted_sum":
            self.kernel_mix = nn.Sequential(
                nn.Linear(bottleneck * n_kernels, n_kernels),
                nn.Softmax(dim=-1),
            )
        else:
            self.kernel_mix = None

        if fusion == "concat_fusion":
            self.merge = nn.Conv1d(
                bottleneck,
                bottleneck,
                kernel_size=self.merge_kernel,
                padding=(self.merge_kernel - 1) // 2,
                groups=bottleneck,
            )
        else:
            self.merge = None

        self.up = nn.Linear(bottleneck, d_model)
        self.drop = nn.Dropout(dropout)
        nn.init.zeros_(self.up.weight)
        nn.init.zeros_(self.up.bias)

    def _fuse(self, branches: list[torch.Tensor]) -> torch.Tensor:
        if self.fusion in {"sum", "weighted_sum"}:
            stacked = torch.stack(branches, dim=-2)
            if self.kernel_mix is not None:
                weights = self.kernel_mix(torch.cat(branches, dim=-1))
                stacked = weights.unsqueeze(-1) * stacked
            return stacked.sum(dim=-2)
        fused = torch.cat(branches, dim=-1)
        if self.merge is not None:
            extra = self.merge(fused.transpose(1, 2)).transpose(1, 2)
            fused = fused + extra
        return fused

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        # hidden_states: (batch, time, dim)
        residual = hidden_states
        hidden = self.act(self.down(self.norm(hidden_states)))
        left, right = hidden.chunk(2, dim=-1)
        right = self.gate_norm(right).transpose(1, 2)
        branches = [conv(right).transpose(1, 2) for conv in self.convs]
        hidden = left * self._fuse(branches)
        return residual + self.drop(self.up(hidden))


class EncoderBlockWithConvAdapter(nn.Module):
    def __init__(self, block: nn.Module, adapter: MultiConvAdapter):
        super().__init__()
        self.block = block
        self.conv_adapter = adapter

    def forward(self, *args, **kwargs):
        hidden_states = self.block(*args, **kwargs)
        if isinstance(hidden_states, tuple):
            return (self.conv_adapter(hidden_states[0]),) + hidden_states[1:]
        return self.conv_adapter(hidden_states)


def get_encoder(model: nn.Module) -> nn.Module:
    core = model.get_base_model() if hasattr(model, "get_base_model") else model
    if hasattr(core, "model") and hasattr(core.model, "encoder"):
        return core.model.encoder
    if hasattr(core, "encoder"):
        return core.encoder
    raise AttributeError("Could not find a Conformer / Parakeet encoder on this model")


def attach_multiconv_adapters(
    model: nn.Module,
    *,
    bottleneck: int,
    kernels: Iterable[int],
    dropout: float,
    skip_bottom_frac: float,
    fusion: str = DEFAULT_FUSION,
    merge_kernel: int = DEFAULT_MERGE_KERNEL,
) -> dict:
    encoder = get_encoder(model)
    layers = encoder.layers
    n_layers = len(layers)
    start = int(n_layers * skip_bottom_frac)
    d_model = int(getattr(encoder.config, "hidden_size", 1280))
    attached = []
    for idx in range(start, n_layers):
        block = layers[idx]
        if isinstance(block, EncoderBlockWithConvAdapter):
            continue
        adapter = MultiConvAdapter(
            d_model,
            bottleneck=bottleneck,
            kernels=kernels,
            dropout=dropout,
            fusion=fusion,
            merge_kernel=merge_kernel,
        )
        try:
            ref = next(block.parameters())
            adapter.to(device=ref.device, dtype=ref.dtype)
        except StopIteration:
            pass
        layers[idx] = EncoderBlockWithConvAdapter(block, adapter)
        attached.append(idx)
    return {
        "n_layers": n_layers,
        "start_layer": start,
        "attached_layers": attached,
        "d_model": d_model,
        "bottleneck": bottleneck,
        "kernels": list(kernels),
        "fusion": fusion,
        "merge_kernel": merge_kernel,
        "dropout": dropout,
    }


def _in_decoder(name: str) -> bool:
    dotted = f".{name}."
    return ".decoder." in dotted or name.startswith("decoder.")


def decoder_lora_targets(model: nn.Module, kinds: Iterable[str]) -> list[str]:
    """Only the 8-layer decoder (self-attn + cross-attn). Never the encoder."""
    kinds = tuple(kinds)
    names = []
    for name, _module in model.named_modules():
        leaf = name.rsplit(".", 1)[-1]
        if leaf not in kinds:
            continue
        if not _in_decoder(name):
            continue
        if "self_attn" not in name and "encoder_attn" not in name:
            continue
        names.append(name)
    if not names:
        raise RuntimeError(
            "No decoder attention projections found. "
            "Is this CohereAsrForConditionalGeneration?"
        )
    return names


def encoder_lora_targets(model: nn.Module, kinds: Iterable[str]) -> list[str]:
    """Conformer self-attn q/k/v/o. Never decoder, never MLP."""
    kinds = tuple(kinds)
    names = []
    for name, _module in model.named_modules():
        leaf = name.rsplit(".", 1)[-1]
        if leaf not in kinds:
            continue
        if _in_decoder(name):
            continue
        if "self_attn" not in name:
            continue
        names.append(name)
    if not names:
        raise RuntimeError("No encoder attention projections found.")
    return names


def lora_targets_for_scope(model: nn.Module, scope: str, kinds: Iterable[str]) -> list[str]:
    if scope == "decoder":
        return decoder_lora_targets(model, kinds)
    if scope == "encoder":
        return encoder_lora_targets(model, kinds)
    if scope == "full":
        return decoder_lora_targets(model, kinds) + encoder_lora_targets(model, kinds)
    raise ValueError(f"Unknown lora scope={scope}")