File size: 2,282 Bytes
d53adc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
# Exported for HuggingFace trust_remote_code loading.
# This file is intentionally self-contained.

from __future__ import annotations

import torch
import torch.nn.functional as functional
from torch import nn



def infer_module_tensor_device(
    module: nn.Module,
    *,
    default_device: torch.device,
) -> torch.device:
    for parameter in module.parameters():
        return parameter.device
    for buffer in module.buffers():
        if isinstance(buffer, torch.Tensor):
            return buffer.device
    return default_device


def apply_preserving_complex_buffers(
    module: nn.Module,
    fn,
    *,
    buffer_names: tuple[str, ...],
    apply_super,
):
    preserved: dict[str, torch.Tensor] = {}
    preserved_device = torch.device("cpu")
    for name in buffer_names:
        tensor = module._buffers.get(name)
        if isinstance(tensor, torch.Tensor) and tensor.is_complex():
            preserved[name] = tensor
            preserved_device = tensor.device
            module._buffers[name] = None
    try:
        result = apply_super()
    finally:
        if preserved:
            target_device = infer_module_tensor_device(
                module,
                default_device=preserved_device,
            )
            for name, tensor in preserved.items():
                module._buffers[name] = tensor.to(device=target_device)
    return result

class RMSNorm(nn.Module):
    """Root Mean Square Layer Normalization."""

    def __init__(self, dim: int, eps: float = 1e-6):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32))
        self.weight._no_weight_decay = True  # type: ignore[attr-defined]

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return functional.rms_norm(
            x,
            (int(x.size(-1)),),
            self.weight,
            float(self.eps),
        )


class StandardLogitMixer(nn.Module):
    """Decoder logits path: final norm followed by output projection."""

    def forward(
        self,
        x: torch.Tensor,
        *,
        norm: RMSNorm,
        output: nn.Module,
    ) -> torch.Tensor:
        return output(norm(x))