""" 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[""], self.tok._vocab[""] 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=1.0, top_k=40, repetition_penalty=1.3, prompt="", greedy=False, top_p=0.9): """Yield the growing continuation base-by-base (for live streaming UIs). greedy=True takes the argmax 6-mer each step (deterministic — the model's single best guess, same decoding as the sequence-recovery metric; the coding domains produce in-frame ATG-start sequences, the low-complexity domains collapse to homopolymers). Sampling (greedy=False) trades that for variety; repetition_penalty then discourages the repeat loops these small specialists fall into.""" m = self.models[domain] # frame-align + cap: our 6-mer tokens tile cleanly only when the context length is a # multiple of 6 (else the model generates out of phase). Trim the leading remainder and # cap to the context window — this is what makes generation match the offline tests. p = clean(prompt) if prompt else "" p = p[-1020:] # 1020 = 170*6, within the 1024-token window p = p[len(p) % 6:] # drop leading remainder so the 6-mers align to the end ids = [self.bos] + (self.tok.encode(p, add_special_tokens=False) if p else []) t = torch.tensor([ids], device=self.dev) bases, emitted = [], [] while sum(len(b) for b in bases) < length: logits = m(input_ids=t).logits[:, -1, :].float() logits[:, :4] = -1e9 # never emit specials # Repetition control applies to BOTH decoders. Greedy without this falls into a # self-reinforcing loop (argmax keeps re-picking the same 6-mer); the penalty plus # a hard block on the last few emitted tokens forces it onward to its next-best, # non-repeating guess — which is what actually de-degenerates the output. if greedy: # argmax with repeat penalty + a hard block on recent tokens (greedy alone loops) if repetition_penalty and repetition_penalty != 1.0: for tid in set(emitted[-12:]): v = logits[0, tid] logits[0, tid] = v / repetition_penalty if v > 0 else v * repetition_penalty for tid in set(emitted[-8:]): logits[0, tid] = -1e9 ti = int(logits.argmax()) else: # nucleus (top-p) sampling — the same decoder Carbon uses. Adapts the candidate # set to the distribution (keeps only tokens carrying mass), which escapes the # low-complexity GC/AT loops that a fixed top-k falls into. No repetition penalty. logits = logits / max(temperature, 1e-6) sl, si = torch.sort(logits[0], descending=True) cum = torch.cumsum(F.softmax(sl, dim=-1), dim=-1) rm = cum > top_p rm[1:] = rm[:-1].clone(); rm[0] = False logits[0, si[rm]] = -1e9 ti = int(torch.multinomial(F.softmax(logits[0], dim=-1), 1)) emitted.append(ti) t = torch.cat([t, torch.tensor([[ti]], device=self.dev)], dim=1) bases.append(self.tok._ids_to_tokens[ti]) yield "".join(bases)[:length] def generate(self, domain, length=180, temperature=1.0, top_k=40, repetition_penalty=1.3, prompt="", greedy=False): out = "" for out in self.generate_stream(domain, length, temperature, top_k, repetition_penalty, prompt, greedy): pass return out