File size: 5,397 Bytes
947f8cf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
daisychain.py -- self-contained inference for the DaisyChain genomic modular mind.

4 dense ~74M DNA/RNA specialists (eukaryote, prokaryote, mrna, mrna_splice), each
per-domain-distilled from Carbon-500M, behind a learned router (MLP on PCA(hidden)
+ per-specialist surprise). route() picks the home specialist; generate() / surprise()
expose the rest. No training/datasets dependency -- only model.py, specialist_presets.py,
spike_tokenizer.py, registry.py + the bundled tokenizer.json / *.safetensors / router2.pt.
"""
from __future__ import annotations
import os, math
import torch
import torch.nn.functional as F

HERE = os.path.dirname(os.path.abspath(__file__))
from model import SpikeWhaleLM
from specialist_presets import generic_specialist_config
from spike_tokenizer import SpikeTokenizer
import registry

TOK_JSON = os.path.join(HERE, "tokenizer.json")
_TRANS = str.maketrans({"U": "T", "u": "T", "a": "A", "c": "C", "g": "G", "t": "T"})
LN2 = math.log(2)


def clean(seq: str) -> str:
    seq = seq.translate(_TRANS).upper()
    return "".join(c if c in "ACGT" else "N" for c in seq)


class _RouterMLP(torch.nn.Module):
    def __init__(self, dim, h=64):
        super().__init__()
        self.net = torch.nn.Sequential(torch.nn.Linear(dim, h), torch.nn.ReLU(),
                                       torch.nn.Dropout(0.0), torch.nn.Linear(h, 4))
    def forward(self, x): return self.net(x)


class DaisyChain:
    DESCRIPTIONS = {
        "eukaryote": "Eukaryotic genomic DNA",
        "prokaryote": "Bacterial / prokaryotic DNA",
        "mrna": "Mature mRNA (coding transcript)",
        "mrna_splice": "Pre-mRNA / splice-site regions",
    }

    def __init__(self, root=HERE, device="cpu"):
        self.dev = device
        self.tok = SpikeTokenizer(vocab_file=os.path.join(root, "tokenizer.json"))
        self.bos, self.eos = self.tok._vocab["<bos>"], self.tok._vocab["<eos>"]
        self.models = {}
        from safetensors.torch import load_file
        for d in registry.ACTIVE:
            ckpt = os.path.join(root, d, "model.safetensors")
            if not os.path.exists(ckpt):
                continue
            cfg = generic_specialist_config(self.tok.vocab_size, position=registry.spec(d)["position"])
            m = SpikeWhaleLM(cfg).to(device).eval()
            sd = load_file(ckpt, device=device)
            m.load_state_dict({k: (v.float() if v.is_floating_point() else v) for k, v in sd.items()})
            for p in m.parameters():
                p.requires_grad_(False)
            self.models[d] = m
        self.domains = list(self.models)
        self.router2 = None
        r2 = os.path.join(root, "router2.pt")
        if os.path.exists(r2):
            d = torch.load(r2, map_location="cpu")
            if all(x in self.models for x in d["domains"]):
                mlp = _RouterMLP(d["k"] + 4, d["h"]); mlp.load_state_dict(d["mlp"]); mlp.eval()
                d["net"] = mlp; self.router2 = d

    @torch.no_grad()
    def _scores_hidden(self, seq):
        ids = [self.bos] + self.tok.encode(clean(seq), add_special_tokens=False) + [self.eos]
        t = torch.tensor([ids], device=self.dev)
        scores, hids = {}, {}
        for d, m in self.models.items():
            hids[d] = m.model(input_ids=t)[0][0].mean(0)
            scores[d] = float(m(input_ids=t, labels=t).loss)
        return scores, hids

    def surprise(self, seq):
        """Per-specialist bits/base (lower = more 'at home')."""
        s, _ = self._scores_hidden(seq)
        return {d: s[d] / 6 / LN2 for d in self.domains}

    @torch.no_grad()
    def route(self, seq):
        """Return (home_domain, bits_per_base_dict). Uses the learned MLP router."""
        scores, hids = self._scores_hidden(seq)
        bpb = {d: scores[d] / 6 / LN2 for d in self.domains}
        if self.router2 is not None:
            r = self.router2
            hidden = torch.cat([hids[d] for d in r["domains"]])
            bits = torch.tensor([scores[d] for d in r["domains"]])
            z = (hidden - r["pca_mu"]) @ r["P"]
            feat = ((torch.cat([z, bits]) - r["mu"]) / r["sd"])
            best = r["domains"][int(r["net"](feat.unsqueeze(0)).argmax(1))]
        else:
            best = min(scores, key=scores.get)
        return best, bpb

    @torch.no_grad()
    def generate_stream(self, domain, length=180, temperature=0.9, top_k=20, prompt=""):
        """Yield the growing continuation base-by-base (for live streaming UIs)."""
        m = self.models[domain]
        ids = [self.bos] + (self.tok.encode(clean(prompt), add_special_tokens=False) if prompt else [])
        t = torch.tensor([ids], device=self.dev)
        bases = []
        while sum(len(b) for b in bases) < length:
            logits = m(input_ids=t).logits[:, -1, :] / max(temperature, 1e-6)
            logits[:, :4] = -1e9
            if top_k > 0:
                v, _ = torch.topk(logits, top_k)
                logits[logits < v[:, [-1]]] = -1e9
            nxt = torch.multinomial(F.softmax(logits, dim=-1), 1)
            t = torch.cat([t, nxt], dim=1)
            bases.append(self.tok._ids_to_tokens[int(nxt)])
            yield "".join(bases)[:length]

    def generate(self, domain, length=180, temperature=0.9, top_k=20, prompt=""):
        out = ""
        for out in self.generate_stream(domain, length, temperature, top_k, prompt):
            pass
        return out