File size: 5,568 Bytes
75ee161
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Phase C — Trainable semantic encoder (cortex learning).
Hashed n-gram features (fixed, biological 'receptive fields') + learned
projection W trained with InfoNCE contrastive loss (NumPy, sparse updates).

Baseline for comparison: lhc_core.Encoder (pure random-indexing hashing).
"""
import hashlib
import re

import numpy as np


def tokenize(text):
    return re.findall(r"\w+", text.lower(), flags=re.UNICODE)


def hashed_features(text, F=2 ** 15):
    """Hashing trick over word unigrams + bigrams + char 3-grams (handles
    Arabic morphology and typos). Returns (indices, signs) sparse vector."""
    toks = tokenize(text)
    feats = []
    for i, t in enumerate(toks):
        feats.append("u:" + t)
        if i + 1 < len(toks):
            feats.append("b:" + t + "_" + toks[i + 1])
        w = "#" + t + "#"
        for j in range(len(w) - 2):
            feats.append("c:" + w[j:j + 3])
    acc = {}
    for f in feats:
        h = int(hashlib.md5(f.encode()).hexdigest()[:8], 16)
        i = h % F
        s = 1.0 if (h >> 16) & 1 else -1.0
        acc[i] = acc.get(i, 0.0) + s
    idx = np.array([i for i, v in acc.items() if v != 0], dtype=np.int64)
    val = np.array([v for v in acc.values() if v != 0], dtype=np.float32)
    return idx, val


class TrainableEncoder:
    """e(text) = normalize( x(text) @ W ).  W learned via InfoNCE so that
    augmented views of the same memory land close (semantic retrieval)."""

    def __init__(self, F=2 ** 15, d=128, tau=0.1, seed=0, W=None, features_fn=None):
        self.F, self.d, self.tau = F, d, tau
        self.feat = features_fn or hashed_features   # swappable featurizer (P2)
        rng = np.random.RandomState(seed)
        self.W = W if W is not None else (rng.randn(F, d) * 0.02).astype(np.float32)

    # ---- forward -------------------------------------------------
    def encode(self, text):
        idx, val = self.feat(text, self.F)
        if len(idx) == 0:
            return np.zeros(self.d, dtype=np.float32)
        z = val @ self.W[idx]
        n = float(np.linalg.norm(z))
        return (z / n).astype(np.float32) if n > 0 else z.astype(np.float32)

    def encode_batch(self, texts):
        return np.stack([self.encode(t) for t in texts], axis=0)

    # ---- InfoNCE training (sparse SGD) ---------------------------
    def _train_views(self, view_fn, pool, steps, batch, lr, seed, log_every):
        """view_fn(pool, rng) -> (anchor_text, positive_text)."""
        rng = np.random.RandomState(seed)
        losses = []
        for step in range(steps):
            picks = rng.randint(0, len(pool), batch)
            views = [view_fn(pool[i], rng) for i in picks]
            anchors = [v[0] for v in views]
            posits = [v[1] for v in views]

            feats_a = [self.feat(t, self.F) for t in anchors]
            feats_p = [self.feat(t, self.F) for t in posits]

            def fwd(feats):
                Z = np.zeros((batch, self.d), dtype=np.float32)
                for r, (idx, val) in enumerate(feats):
                    if len(idx):
                        Z[r] = val @ self.W[idx]
                norms = np.linalg.norm(Z, axis=1, keepdims=True)
                E = Z / np.maximum(norms, 1e-8)
                return E, norms

            Ea, na = fwd(feats_a)
            Ep, np_ = fwd(feats_p)

            S = (Ea @ Ep.T) / self.tau                       # B×B
            S -= S.max(axis=1, keepdims=True)
            P = np.exp(S)
            P /= P.sum(axis=1, keepdims=True)
            loss = -float(np.mean(np.log(np.maximum(P[np.arange(batch), np.arange(batch)], 1e-12))))
            losses.append(loss)

            dS = P.copy()
            dS[np.arange(batch), np.arange(batch)] -= 1.0
            dS /= batch
            dEa = (dS @ Ep) / self.tau
            dEp = (dS.T @ Ea) / self.tau

            def bwd(feats, E, norms, dE):
                dots = np.sum(E * dE, axis=1, keepdims=True)
                dZ = (dE - E * dots) / np.maximum(norms, 1e-8)
                for r, (idx, val) in enumerate(feats):
                    if len(idx):
                        self.W[idx] -= lr * np.outer(val.astype(np.float32), dZ[r]).astype(np.float32)

            bwd(feats_a, Ea, na, dEa)
            bwd(feats_p, Ep, np_, dEp)

            if log_every and (step + 1) % log_every == 0:
                lr *= 0.95
        return {"loss_first": float(np.mean(losses[:50])),
                "loss_last": float(np.mean(losses[-50:])),
                "steps": steps}

    def train(self, texts_pool, augment_fn, steps=1500, batch=32, lr=0.05,
              seed=1, log_every=300):
        """Same-language contrastive: two augmented views of one prompt."""
        return self._train_views(
            lambda t, rng: (augment_fn(t, rng), augment_fn(t, rng)),
            texts_pool, steps, batch, lr, seed, log_every)

    def train_pairs(self, pairs, augment_fn, steps=1500, batch=32, lr=0.05,
                    seed=1, log_every=300):
        """Cross-lingual contrastive: (text_a, text_b) = two languages of ONE
        meaning — pulls e.g. English and Arabic views of the same question close."""
        return self._train_views(
            lambda pair, rng: (augment_fn(pair[0], rng), augment_fn(pair[1], rng)),
            pairs, steps, batch, lr, seed, log_every)

    def save(self, path):
        np.savez_compressed(path, W=self.W, F=self.F, d=self.d, tau=self.tau)

    @classmethod
    def load(cls, path):
        z = np.load(path)
        return cls(F=int(z["F"]), d=int(z["d"]), tau=float(z["tau"]), W=z["W"])