File size: 3,505 Bytes
7c5e40e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Position extensions that preserve every frozen source-context lookup exactly."""

from __future__ import annotations

import torch
from torch import nn


class SegmentFactorizedPositionExtension(nn.Module):
    """Use frozen absolute rows in-range and segment/offset factors out-of-range.

    The accepted source checkpoint is an exact branch for positions below
    ``source_context``.  Only positions beyond that boundary use the new
    factorized parameters, so the extension cannot silently alter the frozen
    behavior it is intended to extend.
    """

    def __init__(
        self,
        source: nn.Embedding,
        *,
        target_context: int,
        segment_size: int,
    ) -> None:
        super().__init__()
        if target_context <= source.num_embeddings:
            raise ValueError("target context must exceed the source position table")
        if segment_size <= 0 or source.num_embeddings % segment_size:
            raise ValueError("segment size must divide the source position table")
        if target_context % segment_size:
            raise ValueError("segment size must divide target context")
        self.source_context = int(source.num_embeddings)
        self.target_context = int(target_context)
        self.segment_size = int(segment_size)
        self.embedding_dim = int(source.embedding_dim)
        self.frozen_source = nn.Embedding.from_pretrained(
            source.weight.detach().float().clone(), freeze=True,
        )

        table = self.frozen_source.weight.view(-1, segment_size, self.embedding_dim)
        local = table.mean(dim=0)
        source_age = (table - local.unsqueeze(0)).mean(dim=1)
        future_segments = target_context // segment_size - table.shape[0]
        future = torch.stack([
            source_age[index % source_age.shape[0]]
            for index in range(future_segments)
        ])
        self.local_offsets = nn.Parameter(local.clone())
        self.future_segment_age = nn.Parameter(future.clone())

    @property
    def num_embeddings(self) -> int:
        return self.target_context

    @property
    def weight(self) -> torch.Tensor:
        positions = torch.arange(self.target_context, device=self.local_offsets.device)
        return self(positions)

    def forward(self, positions: torch.Tensor) -> torch.Tensor:
        if positions.numel() and (
            int(positions.min()) < 0 or int(positions.max()) >= self.target_context
        ):
            raise IndexError("position index is outside the extended context")
        positions = positions.to(dtype=torch.long)
        source_position = positions.clamp_max(self.source_context - 1)
        source = self.frozen_source(source_position)
        future = positions >= self.source_context
        if not bool(future.any()):
            return source
        local_index = positions.remainder(self.segment_size)
        segment_index = torch.div(
            positions - self.source_context, self.segment_size, rounding_mode="floor"
        ).clamp_min(0)
        extension = self.local_offsets[local_index] + self.future_segment_age[segment_index]
        return torch.where(future.unsqueeze(-1), extension, source)

    def first_rows_exact(self, source: torch.Tensor) -> bool:
        positions = torch.arange(self.source_context, device=self.local_offsets.device)
        observed = self(positions).detach().cpu()
        return torch.equal(observed, source.detach().float().cpu())


__all__ = ["SegmentFactorizedPositionExtension"]