File size: 17,794 Bytes
24af195
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
"""
Memory Sparse Attention (MSA) for PC-SHO-DLM

Implements the MSA framework from "Memory Sparse Attention for Efficient
End-to-End Memory Model Scaling to 100M Tokens" (Chen et al., 2025),
integrated with predictive-coding energy minimization.

Key components:
1. Router Projectors (W_QR, W_KR) for document-level relevance scoring
2. Chunk-wise KV compression via mean pooling
3. Top-k sparse document selection
4. Document-wise RoPE for extrapolation to 100M tokens
5. Memory Interleave for multi-hop reasoning via iterative settling
6. Contrastive auxiliary loss for router training

The deep integration with PC-SHO-DLM:
- Router scoring is part of the energy function (settling optimizes retrieval)
- Precision heads and routers share information (uncertainty = retrieval need)
- Unified mode: retrieval improves during settling as parameters update
"""

import math
from dataclasses import dataclass, field
from typing import Optional, Tuple, List

import torch
import torch.nn as nn
import torch.nn.functional as F


@dataclass
class MSAConfig:
    """Configuration for Memory Sparse Attention."""
    chunk_size: int = 64          # tokens per chunk for compression
    top_k: int = 16               # number of documents to retrieve
    router_dim: int = 128         # dimension of router projections
    n_router_heads: int = 8       # number of router heads
    apply_from_layer: int = 6     # only apply MSA to upper half of layers (MSA finding)
    aux_loss_weight: float = 0.1  # weight of contrastive routing loss
    aux_temperature: float = 0.05 # temperature for contrastive loss
    rope_base: float = 10000.0    # RoPE base frequency


class DocumentWiseRoPE(nn.Module):
    """Document-wise Rotary Position Embedding.

    Each document gets independent position IDs starting from 0.
    Query tokens get global offset by top_k.
    This enables training on 64k but extrapolating to 100M tokens.
    """

    def __init__(self, d_model: int, max_len: int = 8192, base: float = 10000.0):
        super().__init__()
        self.d_model = d_model
        self.max_len = max_len
        self.base = base

        # Precompute frequencies
        inv_freq = 1.0 / (base ** (torch.arange(0, d_model, 2).float() / d_model))
        self.register_buffer("inv_freq", inv_freq)

    def _compute_rope(self, positions: torch.Tensor, dim: int) -> Tuple[torch.Tensor, torch.Tensor]:
        """Compute cos and sin for given position indices."""
        # positions: (B, S) or (S,)
        freqs = torch.einsum("...s,d->...sd", positions.float(), self.inv_freq[:dim // 2].to(positions.device))
        cos = freqs.cos()
        sin = freqs.sin()
        return cos, sin

    def apply_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
        """Apply rotary embeddings to input tensor."""
        # x: (..., S, D), cos/sin: (..., S, D//2)
        d = x.shape[-1]
        x1, x2 = x[..., :d // 2], x[..., d // 2:]
        return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)

    def forward(self, x: torch.Tensor, doc_boundaries: Optional[torch.Tensor] = None,
                global_offset: int = 0) -> torch.Tensor:
        """Apply document-wise RoPE.

        Args:
            x: (B, S, D) input
            doc_boundaries: (B, S) tensor of document IDs per position.
                           Positions within the same doc get local IDs.
                           If None, standard positional encoding.
            global_offset: offset for query positions (= top_k retrieved docs)
        """
        B, S, D = x.shape

        if doc_boundaries is not None:
            # Document-wise: each doc starts at position 0
            positions = torch.zeros(B, S, device=x.device, dtype=torch.long)
            for b in range(B):
                for doc_id in doc_boundaries[b].unique():
                    doc_mask = doc_boundaries[b] == doc_id
                    positions[b, doc_mask] = torch.arange(doc_mask.sum(), device=x.device)
        else:
            # Standard positional encoding with optional offset
            positions = torch.arange(S, device=x.device).unsqueeze(0).expand(B, -1) + global_offset

        cos, sin = self._compute_rope(positions, D)
        return self.apply_rope(x, cos, sin)


class RouterProjector(nn.Module):
    """Learned router for document relevance scoring.

    Separate from the main K/Q projectors — dedicated to retrieval.
    Produces routing keys (KR) and routing queries (QR) in a shared space.
    """

    def __init__(self, d_model: int, router_dim: int, n_heads: int):
        super().__init__()
        self.n_heads = n_heads
        self.d_head = router_dim // n_heads
        self.q_proj = nn.Linear(d_model, router_dim)
        self.k_proj = nn.Linear(d_model, router_dim)

    def project_query(self, h_q: torch.Tensor) -> torch.Tensor:
        """Project query hidden states to routing space. (B, S, router_dim)"""
        return self.q_proj(h_q)

    def project_key(self, h_doc: torch.Tensor) -> torch.Tensor:
        """Project document hidden states to routing space. (B, S, router_dim)"""
        return self.k_proj(h_doc)


class MemoryBank:
    """Stores compressed KV representations of documents.

    Offline: encode documents → chunk → mean-pool K,V,KR → store
    Online: query → router scores → top-k → load compressed KV

    The bank stores three things per document per layer:
    - K_bar: compressed keys (n_chunks, n_heads, d_head)
    - V_bar: compressed values (n_chunks, n_heads, d_head)
    - KR_bar: compressed routing keys (n_chunks, router_dim)
    """

    def __init__(self, chunk_size: int = 64):
        self.chunk_size = chunk_size
        self.documents = {}  # doc_id -> {layer_id -> {K_bar, V_bar, KR_bar}}
        self.doc_ids = []

    def add_document(self, doc_id: str, layer_kvs: dict):
        """Add a document's compressed KV to the bank.

        Args:
            doc_id: unique identifier
            layer_kvs: {layer_idx: {"K": tensor, "V": tensor, "KR": tensor}}
                       Each tensor has shape (n_chunks, ...)
        """
        self.documents[doc_id] = layer_kvs
        if doc_id not in self.doc_ids:
            self.doc_ids.append(doc_id)

    def get_routing_keys(self, layer_idx: int) -> Tuple[torch.Tensor, List[str]]:
        """Get all routing keys for a layer. Returns (N_total_chunks, router_dim) + doc IDs."""
        keys = []
        ids = []
        for doc_id in self.doc_ids:
            if layer_idx in self.documents[doc_id]:
                kr = self.documents[doc_id][layer_idx]["KR"]
                keys.append(kr)
                ids.extend([doc_id] * kr.shape[0])
        if keys:
            return torch.cat(keys, dim=0), ids
        return None, []

    def get_kv(self, doc_ids: List[str], layer_idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
        """Get compressed K,V for selected documents at a layer."""
        ks, vs = [], []
        for did in doc_ids:
            if did in self.documents and layer_idx in self.documents[did]:
                ks.append(self.documents[did][layer_idx]["K"])
                vs.append(self.documents[did][layer_idx]["V"])
        if ks:
            return torch.cat(ks, dim=0), torch.cat(vs, dim=0)
        return None, None

    def __len__(self):
        return len(self.doc_ids)


def chunk_mean_pool(x: torch.Tensor, chunk_size: int) -> torch.Tensor:
    """Compress sequence via chunk-wise mean pooling.

    Args:
        x: (B, S, D) or (S, D)
        chunk_size: tokens per chunk

    Returns:
        (B, n_chunks, D) or (n_chunks, D)
    """
    if x.dim() == 2:
        S, D = x.shape
        n_chunks = math.ceil(S / chunk_size)
        # Pad to multiple of chunk_size
        if S % chunk_size != 0:
            pad = chunk_size - (S % chunk_size)
            x = F.pad(x, (0, 0, 0, pad))
        return x.view(n_chunks, chunk_size, D).mean(dim=1)
    else:
        B, S, D = x.shape
        n_chunks = math.ceil(S / chunk_size)
        if S % chunk_size != 0:
            pad = chunk_size - (S % chunk_size)
            x = F.pad(x, (0, 0, 0, pad))
        return x.view(B, n_chunks, chunk_size, D).mean(dim=2)


class MSALayer(nn.Module):
    """Memory Sparse Attention layer.

    Replaces standard self-attention with sparse document-level retrieval.
    Applied only to upper layers (lower layers use standard attention).

    The key integration with PC-SHO-DLM:
    - Router scores become part of the energy function
    - Settling optimizes both hidden states AND retrieval quality
    - Precision heads inform router confidence
    """

    def __init__(self, d_model: int, n_heads: int, d_ff: int,
                 msa_config: MSAConfig, dropout: float = 0.1):
        super().__init__()
        self.d_model = d_model
        self.n_heads = n_heads
        self.d_head = d_model // n_heads
        self.msa_config = msa_config

        # Standard Q/K/V projectors
        self.q_proj = nn.Linear(d_model, d_model)
        self.k_proj = nn.Linear(d_model, d_model)
        self.v_proj = nn.Linear(d_model, d_model)
        self.o_proj = nn.Linear(d_model, d_model)

        # Router projectors (separate from main attention)
        self.router = RouterProjector(d_model, msa_config.router_dim, msa_config.n_router_heads)

        # Document-wise RoPE
        self.rope = DocumentWiseRoPE(self.d_head)

        # FFN
        self.ff = nn.Sequential(
            nn.Linear(d_model, d_ff), nn.GELU(),
            nn.Linear(d_ff, d_model), nn.Dropout(dropout),
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)

    def compute_routing_scores(self, h_query: torch.Tensor,
                                memory_routing_keys: torch.Tensor) -> torch.Tensor:
        """Score documents by relevance to query.

        Args:
            h_query: (B, S_q, D) query hidden states
            memory_routing_keys: (N_chunks, router_dim) compressed routing keys

        Returns:
            scores: (B, N_chunks) relevance scores
        """
        # Project query to routing space
        qr = self.router.project_query(h_query)  # (B, S_q, router_dim)

        # Normalize for cosine similarity
        qr_norm = F.normalize(qr, dim=-1)
        kr_norm = F.normalize(memory_routing_keys, dim=-1)

        # Score: max over query tokens, mean over heads
        # (B, S_q, router_dim) @ (N_chunks, router_dim)^T → (B, S_q, N_chunks)
        sim = torch.einsum("bsd,nd->bsn", qr_norm, kr_norm.to(qr_norm.device))

        # Max-pool over query tokens
        scores = sim.max(dim=1).values  # (B, N_chunks)

        return scores

    def sparse_attention(self, h_query: torch.Tensor,
                         memory_k: torch.Tensor, memory_v: torch.Tensor,
                         doc_boundaries: Optional[torch.Tensor] = None) -> torch.Tensor:
        """Attend to query + selected memory KV.

        Args:
            h_query: (B, S_q, D) query hidden states
            memory_k: (B, S_mem, D) compressed keys from selected documents
            memory_v: (B, S_mem, D) compressed values from selected documents

        Returns:
            output: (B, S_q, D)
        """
        B, S_q, D = h_query.shape
        H, d = self.n_heads, self.d_head

        # Project query
        Q = self.q_proj(h_query).view(B, S_q, H, d).transpose(1, 2)

        if memory_k is not None:
            S_mem = memory_k.shape[1]
            # Concatenate memory + query KV
            K_mem = self.k_proj(memory_k).view(B, S_mem, H, d).transpose(1, 2)
            V_mem = self.v_proj(memory_v).view(B, S_mem, H, d).transpose(1, 2)

            K_q = self.k_proj(h_query).view(B, S_q, H, d).transpose(1, 2)
            V_q = self.v_proj(h_query).view(B, S_q, H, d).transpose(1, 2)

            # Apply document-wise RoPE to memory, global RoPE to query
            # (simplified: just offset query positions)
            K_ctx = torch.cat([K_mem, K_q], dim=2)  # (B, H, S_mem+S_q, d)
            V_ctx = torch.cat([V_mem, V_q], dim=2)
        else:
            K_ctx = self.k_proj(h_query).view(B, S_q, H, d).transpose(1, 2)
            V_ctx = self.v_proj(h_query).view(B, S_q, H, d).transpose(1, 2)

        # Standard scaled dot-product attention
        scale = math.sqrt(d)
        attn = torch.einsum("bhsd,bhtd->bhst", Q, K_ctx) / scale
        attn = F.softmax(attn, dim=-1)
        attn = self.dropout(attn)

        out = torch.einsum("bhst,bhtd->bhsd", attn, V_ctx)
        out = out.transpose(1, 2).reshape(B, S_q, D)

        return self.o_proj(out)

    def forward(self, x: torch.Tensor,
                memory_k: Optional[torch.Tensor] = None,
                memory_v: Optional[torch.Tensor] = None) -> torch.Tensor:
        """Forward pass with optional memory context.

        If memory_k/v are provided, uses sparse attention over memory + local.
        Otherwise, falls back to standard bidirectional self-attention.
        """
        residual = x
        x = self.norm1(x)
        x = residual + self.dropout(self.sparse_attention(x, memory_k, memory_v))

        residual = x
        x = self.norm2(x)
        x = residual + self.ff(x)

        return x


def compute_routing_aux_loss(scores_pos: torch.Tensor, scores_neg: torch.Tensor,
                              temperature: float = 0.05) -> torch.Tensor:
    """Contrastive auxiliary loss for router training (MSA Eq. 5).

    Pushes positive document scores above negative document scores.

    Args:
        scores_pos: (B, n_pos) scores for relevant documents
        scores_neg: (B, n_neg) scores for irrelevant documents
        temperature: softmax temperature

    Returns:
        loss: scalar
    """
    # For each positive, contrast against all negatives
    # L = -1/|P| * sum_i log(exp(s+_i/Ï„) / (exp(s+_i/Ï„) + sum_j exp(s-_j/Ï„)))
    pos_exp = (scores_pos / temperature).exp()  # (B, n_pos)
    neg_exp_sum = (scores_neg / temperature).exp().sum(dim=-1, keepdim=True)  # (B, 1)

    log_prob = (scores_pos / temperature) - torch.log(pos_exp + neg_exp_sum + 1e-10)
    loss = -log_prob.mean()

    return loss


class MemoryEncoder:
    """Encodes documents into the memory bank.

    Offline process: runs each document through the model,
    extracts K, V, and KR at each MSA layer, compresses via
    chunk-wise mean pooling, and stores in the MemoryBank.
    """

    @staticmethod
    @torch.no_grad()
    def encode_document(model, text: str, doc_id: str, memory_bank: MemoryBank,
                        msa_layers: nn.ModuleList, chunk_size: int = 64,
                        device: str = "cpu") -> None:
        """Encode a single document into the memory bank.

        Args:
            model: PCSHODLM model
            text: raw text of the document
            doc_id: unique identifier
            memory_bank: bank to store compressed KV
            msa_layers: list of MSALayer modules
            chunk_size: compression chunk size
        """
        # Encode text as bytes
        tokens = torch.tensor(
            [min(b + 1, 256) for b in text.encode("utf-8")[:model.config.max_seq_len]],
            dtype=torch.long
        ).unsqueeze(0).to(device)

        # Pad if needed
        if tokens.shape[1] < model.config.max_seq_len:
            tokens = F.pad(tokens, (0, model.config.max_seq_len - tokens.shape[1]))

        # Run through model to get hidden states at each layer
        t = torch.ones(1, dtype=torch.long, device=device)  # dummy timestep
        h = model.embed_input(tokens, t)
        layer_kvs = {}

        n_layers = len(model.forward_blocks)
        msa_start = n_layers // 2  # MSA applies to upper half

        for l, block in enumerate(model.forward_blocks):
            h = block(h)

            # For MSA layers, extract and compress KV + routing keys
            if l >= msa_start:
                msa_idx = l - msa_start
                if msa_idx >= len(msa_layers):
                    continue
                msa_layer = msa_layers[msa_idx]

                # Project to K, V spaces
                K = msa_layer.k_proj(h.squeeze(0))   # (S, D)
                V = msa_layer.v_proj(h.squeeze(0))   # (S, D)
                KR = msa_layer.router.project_key(h.squeeze(0))  # (S, router_dim)

                # Compress via chunk-wise mean pooling
                K_bar = chunk_mean_pool(K, chunk_size)   # (n_chunks, D)
                V_bar = chunk_mean_pool(V, chunk_size)
                KR_bar = chunk_mean_pool(KR, chunk_size)  # (n_chunks, router_dim)

                layer_kvs[l] = {"K": K_bar, "V": V_bar, "KR": KR_bar}

        memory_bank.add_document(doc_id, layer_kvs)


# =============================================================================
# Integration helper: create MSA layers for a PC-SHO-DLM model
# =============================================================================

def create_msa_layers(model_config, msa_config: Optional[MSAConfig] = None) -> nn.ModuleList:
    """Create MSA layers for the upper half of the model.

    Args:
        model_config: PCSHOConfig
        msa_config: MSAConfig (uses defaults if None)

    Returns:
        ModuleList of MSALayer, one per upper-half layer
    """
    if msa_config is None:
        msa_config = MSAConfig()

    n_msa_layers = model_config.n_layers - msa_config.apply_from_layer
    if n_msa_layers <= 0:
        n_msa_layers = model_config.n_layers // 2

    layers = nn.ModuleList([
        MSALayer(
            d_model=model_config.d_model,
            n_heads=model_config.n_heads,
            d_ff=model_config.d_ff,
            msa_config=msa_config,
            dropout=model_config.dropout,
        )
        for _ in range(n_msa_layers)
    ])

    return layers