File size: 4,602 Bytes
f532887
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Encode memory text into a four-token SID; no MemR/MemE weights required."""
import argparse
import json
from pathlib import Path

import numpy as np


class SIDEncoder:
    def __init__(self, codebook_dir=None):
        root = Path(codebook_dir) if codebook_dir else Path(__file__).parent / 'codebook'
        self.config = json.loads((root / 'config.json').read_text())
        self.codebooks = [np.load(root / f'codebook_{i}.npy', allow_pickle=False)
                          for i in range(4)]
        for c, size in zip(self.codebooks, self.config['codebook_sizes']):
            if c.shape != (size, self.config['embedding_dim']) or not np.isfinite(c).all():
                raise ValueError('Invalid codebook')
        self.tokenizer = self.model = None

    def encode_embeddings(self, embeddings):
        """Accept already L2-normalized embedding vectors, NOT arbitrary model vectors."""
        residual = np.asarray(embeddings, dtype=np.float32)
        if residual.ndim == 1:
            residual = residual[None, :]
        if residual.ndim != 2 or residual.shape[1] != self.config['embedding_dim']:
            raise ValueError('Expected shape [N, 1024]')
        if not np.isfinite(residual).all():
            raise ValueError('Embeddings must be finite')
        if not np.allclose(np.linalg.norm(residual, axis=1), 1, atol=0.005):
            raise ValueError('Expected L2-normalized Qwen3 embeddings')
        codes = []
        for centers, weight in zip(self.codebooks, self.config['spherical_weight_per_level']):
            x2 = np.einsum('ij,ij->i', residual, residual)[:, None]
            c2 = np.einsum('ij,ij->i', centers, centers)[None, :]
            dot = residual @ centers.T
            euclidean = np.maximum(x2 + c2 - 2 * dot, 0)
            cosine_distance = np.maximum(1 - dot / (
                np.sqrt(np.maximum(x2, 1e-12)) * np.sqrt(np.maximum(c2, 1e-12))), 0)
            indices = np.argmin((1 - weight) * euclidean + weight * cosine_distance, axis=1)
            codes.append(indices)
            residual = residual - centers[indices]
        return np.stack(codes, axis=1)

    def load_embedding_model(self, device='cpu', model_path=None):
        import torch
        from transformers import AutoModel, AutoTokenizer
        cfg = self.config['embedding']
        kwargs = {} if model_path else {'revision': cfg['revision']}
        name = model_path or cfg['model']
        self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
        self.model = AutoModel.from_pretrained(name, torch_dtype=(
            torch.float16 if str(device).startswith('cuda') else torch.float32), **kwargs).to(device).eval()

    def embed(self, memories):
        import torch
        if self.model is None:
            self.load_embedding_model()
        result = []
        # Single-item encoding avoids the legacy padded-batch last-token ambiguity.
        for memory in memories:
            if not isinstance(memory, str) or not memory.strip():
                raise ValueError('Memory must be nonempty text')
            inputs = self.tokenizer(self.config['embedding']['instruction'] + memory,
                                    return_tensors='pt', truncation=True, max_length=8192)
            inputs = {k: v.to(self.model.device) for k, v in inputs.items()}
            with torch.inference_mode():
                hidden = self.model(**inputs).last_hidden_state
                eos = torch.where(inputs['input_ids'][0] == self.tokenizer.eos_token_id)[0]
                index = int(eos[-1]) if len(eos) else hidden.shape[1] - 1
                vector = torch.nn.functional.normalize(hidden[:, index, :], p=2, dim=1)
            result.append(vector.float().cpu().numpy()[0])
        return np.asarray(result, dtype=np.float32)

    def encode(self, memories):
        if isinstance(memories, str):
            memories = [memories]
        if not memories:
            return []
        codes = self.encode_embeddings(self.embed(memories))
        return [{'sid_codes': row.tolist(), 'sid': ''.join(
            f'<SID_L{i+1}_{int(c)}>' for i, c in enumerate(row))} for row in codes]


if __name__ == '__main__':
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument('--memory', required=True)
    p.add_argument('--device', default='cpu')
    p.add_argument('--embedding-model', help='Optional local Qwen3-Embedding-0.6B snapshot')
    args = p.parse_args()
    encoder = SIDEncoder()
    encoder.load_embedding_model(args.device, args.embedding_model)
    print(json.dumps(encoder.encode(args.memory)[0], ensure_ascii=False))