File size: 2,460 Bytes
a9c0188
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Shared low-level modules for the Inkling MLX port."""

from __future__ import annotations

import mlx.core as mx
import mlx.nn as nn


class RMSNorm(nn.Module):
    """Llama-style RMSNorm (compute in fp32, weight is a gain).

    Matches ``LlamaRMSNorm``: ``x_fp32 * rsqrt(mean(x^2) + eps) * weight``.
    """

    def __init__(self, dims: int, eps: float = 1e-6):
        super().__init__()
        self.weight = mx.ones((dims,))
        self.eps = eps

    def __call__(self, x: mx.array) -> mx.array:
        return mx.fast.rms_norm(x, self.weight, self.eps)


class ShortConvolution(nn.Module):
    """Depthwise causal 1-D convolution with a residual add, computed in fp32.

    Mirrors ``InklingShortConvolution``: a per-channel (groups == channels) causal
    conv1d of ``kernel_size`` taps, no bias, no activation, then ``out + input``.
    The reference keeps this module in fp32 regardless of the model dtype
    (``_keep_in_fp32_modules_strict``), so we upcast here too.

    Weight layout (MLX ``conv1d``): ``[channels, kernel_size, 1]``.
    """

    def __init__(self, channels: int, kernel_size: int):
        super().__init__()
        self.channels = channels
        self.kernel_size = kernel_size
        # [C_out, K, C_in // groups] with groups == channels -> [C, K, 1]
        self.weight = mx.zeros((channels, kernel_size, 1))

    def __call__(self, x: mx.array, mask: mx.array | None = None, cache=None) -> mx.array:
        # x: [batch, seq, channels]
        in_dtype = x.dtype
        xf = x.astype(mx.float32)
        residual = xf
        if mask is not None:
            xf = xf * mask.astype(mx.float32)
        k = self.kernel_size
        B, seq, C = xf.shape
        w = self.weight.astype(mx.float32)
        if cache is not None:
            # left-context = cached last (k-1) inputs (zeros on the first call);
            # a "valid" conv over [left, xf] yields exactly `seq` causal outputs.
            left = cache.state if cache.state is not None else mx.zeros((B, k - 1, C), dtype=mx.float32)
            x_in = mx.concatenate([left, xf], axis=1)
            out = mx.conv1d(x_in, w, padding=0, groups=self.channels)
            cache.state = x_in[:, -(k - 1):, :]
        else:
            # causal: left-pad by (k-1), keep first `seq` outputs (== zero left-context)
            out = mx.conv1d(xf, w, padding=k - 1, groups=self.channels)[:, :seq, :]
        out = out + residual
        return out.astype(in_dtype)