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"]
|