1import pickle, math, numpy as np
2from collections import Counter
3from itertools import product
4
5# Download from HuggingFace
6from huggingface_hub import hf_hub_download
7general_path = hf_hub_download("Seyomi/synthguard-kmer", "general_model.pkl")
8short_path = hf_hub_download("Seyomi/synthguard-kmer", "short_model.pkl")
9
10with open(general_path, "rb") as f: general_model = pickle.load(f)
11with open(short_path, "rb") as f: short_model = pickle.load(f)
12
13# ── Feature extractor (5,533 features — must match exactly) ──────────────────
14CODON_TABLE = {
15 'TTT':'F','TTC':'F','TTA':'L','TTG':'L','CTT':'L','CTC':'L','CTA':'L','CTG':'L',
16 'ATT':'I','ATC':'I','ATA':'I','ATG':'M','GTT':'V','GTC':'V','GTA':'V','GTG':'V',
17 'TCT':'S','TCC':'S','TCA':'S','TCG':'S','CCT':'P','CCC':'P','CCA':'P','CCG':'P',
18 'ACT':'T','ACC':'T','ACA':'T','ACG':'T','GCT':'A','GCC':'A','GCA':'A','GCG':'A',
19 'TAT':'Y','TAC':'Y','TAA':'*','TAG':'*','CAT':'H','CAC':'H','CAA':'Q','CAG':'Q',
20 'AAT':'N','AAC':'N','AAA':'K','AAG':'K','GAT':'D','GAC':'D','GAA':'E','GAG':'E',
21 'TGT':'C','TGC':'C','TGA':'*','TGG':'W','CGT':'R','CGC':'R','CGA':'R','CGG':'R',
22 'AGT':'S','AGC':'S','AGA':'R','AGG':'R','GGT':'G','GGC':'G','GGA':'G','GGG':'G',
23}
24_AA_CODONS = {}
25for c, a in CODON_TABLE.items(): _AA_CODONS.setdefault(a, []).append(c)
26ALL_CODONS = sorted(CODON_TABLE.keys())
27AMINO_ACIDS = sorted(a for a in set(CODON_TABLE.values()) if a != '*')
28VOCAB = {k: ["".join(p) for p in product("ACGT", repeat=k)] for k in [3,4,5,6]}
29
30# Kazusa DB (E. coli / human / yeast) for CAI
31_ECOLI = {'TTT':22.0,'TTC':16.5,'TTA':13.9,'TTG':13.1,'CTT':10.9,'CTC':10.0,'CTA':3.8,'CTG':52.7,'ATT':28.8,'ATC':25.1,'ATA':4.4,'ATG':27.4,'GTT':19.5,'GTC':14.7,'GTA':10.8,'GTG':25.9,'TCT':7.8,'TCC':8.8,'TCA':7.0,'TCG':8.7,'CCT':7.2,'CCC':5.6,'CCA':8.4,'CCG':23.3,'ACT':9.0,'ACC':23.4,'ACA':7.2,'ACG':14.6,'GCT':15.3,'GCC':25.8,'GCA':20.6,'GCG':33.5,'TAT':16.3,'TAC':12.5,'TAA':2.0,'TAG':0.3,'CAT':13.2,'CAC':9.6,'CAA':15.5,'CAG':28.7,'AAT':22.3,'AAC':22.4,'AAA':33.6,'AAG':10.1,'GAT':32.2,'GAC':19.0,'GAA':39.8,'GAG':18.3,'TGT':5.0,'TGC':6.5,'TGA':1.0,'TGG':15.2,'CGT':21.1,'CGC':21.7,'CGA':3.7,'CGG':5.3,'AGT':8.7,'AGC':15.8,'AGA':3.5,'AGG':2.9,'GGT':24.7,'GGC':29.5,'GGA':8.0,'GGG':11.5}
32_HUMAN = {'TTT':17.6,'TTC':20.3,'TTA':7.7,'TTG':12.9,'CTT':13.2,'CTC':19.6,'CTA':7.2,'CTG':39.6,'ATT':16.0,'ATC':20.8,'ATA':7.5,'ATG':22.0,'GTT':11.0,'GTC':14.5,'GTA':7.1,'GTG':28.1,'TCT':15.2,'TCC':17.7,'TCA':12.2,'TCG':4.4,'CCT':17.5,'CCC':19.8,'CCA':16.9,'CCG':6.9,'ACT':13.1,'ACC':18.9,'ACA':15.1,'ACG':6.1,'GCT':18.4,'GCC':27.7,'GCA':15.8,'GCG':7.4,'TAT':12.2,'TAC':15.3,'TAA':1.0,'TAG':0.8,'CAT':10.9,'CAC':15.1,'CAA':12.3,'CAG':34.2,'AAT':17.0,'AAC':19.1,'AAA':24.4,'AAG':31.9,'GAT':21.8,'GAC':25.1,'GAA':29.0,'GAG':39.6,'TGT':10.6,'TGC':12.6,'TGA':1.6,'TGG':13.2,'CGT':4.5,'CGC':10.4,'CGA':6.2,'CGG':11.4,'AGT':15.2,'AGC':19.5,'AGA':11.5,'AGG':11.4,'GGT':10.8,'GGC':22.2,'GGA':16.5,'GGG':16.5}
33_YEAST = {'TTT':26.2,'TTC':18.4,'TTA':26.2,'TTG':27.2,'CTT':12.3,'CTC':5.4,'CTA':13.4,'CTG':10.5,'ATT':30.1,'ATC':17.2,'ATA':17.8,'ATG':20.9,'GTT':22.1,'GTC':11.8,'GTA':11.8,'GTG':10.8,'TCT':23.5,'TCC':14.2,'TCA':18.7,'TCG':8.6,'CCT':13.5,'CCC':6.8,'CCA':18.3,'CCG':5.4,'ACT':20.3,'ACC':13.1,'ACA':17.9,'ACG':8.1,'GCT':21.1,'GCC':12.6,'GCA':16.0,'GCG':6.2,'TAT':18.8,'TAC':14.8,'TAA':1.1,'TAG':0.5,'CAT':13.6,'CAC':7.8,'CAA':27.3,'CAG':12.1,'AAT':35.9,'AAC':24.8,'AAA':41.9,'AAG':30.8,'GAT':37.6,'GAC':20.2,'GAA':45.0,'GAG':19.2,'TGT':8.1,'TGC':4.8,'TGA':0.7,'TGG':10.4,'CGT':6.4,'CGC':2.6,'CGA':3.0,'CGG':1.7,'AGT':14.2,'AGC':9.8,'AGA':21.3,'AGG':9.2,'GGT':23.9,'GGC':9.8,'GGA':10.9,'GGG':6.0}
34
35def _ref_rscu(ft):
36 rscu = {}
37 for aa, codons in _AA_CODONS.items():
38 if aa == '*':
39 for c in codons: rscu[c] = 1.0; continue
40 mf = max(ft.get(c, 0.1) for c in codons)
41 for c in codons: rscu[c] = ft.get(c, 0.1) / mf if mf > 0 else 1.0
42 return rscu
43
44_ECOLI_RSCU = _ref_rscu(_ECOLI)
45_HUMAN_RSCU = _ref_rscu(_HUMAN)
46_YEAST_RSCU = _ref_rscu(_YEAST)
47
48def _codon_features(seq):
49 cc = Counter(seq[i:i+3] for i in range(0, len(seq)-2, 3) if seq[i:i+3] in CODON_TABLE)
50 rscu = {}
51 for aa, codons in _AA_CODONS.items():
52 if aa == '*':
53 for c in codons: rscu[c] = 1.0; continue
54 tot = sum(cc.get(c,0) for c in codons); n = len(codons)
55 exp = tot/n if tot > 0 else 0
56 for c in codons: rscu[c] = cc.get(c,0)/exp if exp > 0 else 1.0
57 rscu_f = [rscu.get(c, 1.0) for c in ALL_CODONS]
58 def cai(ref):
59 s,n = 0.0,0
60 for c,k in cc.items():
61 if CODON_TABLE.get(c,'*') != '*': s += math.log(max(ref.get(c,0.01),1e-6))*k; n+=k
62 return math.exp(s/n) if n else 0.5
63 aa_tot = sum(k for c,k in cc.items() if CODON_TABLE.get(c,'*')!='*')
64 aa_cnt = Counter({CODON_TABLE[c]:k for c,k in cc.items() if CODON_TABLE.get(c,'*')!='*'})
65 return rscu_f + [cai(_ECOLI_RSCU), cai(_HUMAN_RSCU), cai(_YEAST_RSCU)] + \
66 [aa_cnt.get(a,0)/max(aa_tot,1) for a in AMINO_ACIDS]
67
68def extract_features(seq):
69 seq = seq.upper().replace("U","T")
70 n = max(len(seq),1); cnt = Counter(seq); tot = sum(cnt.values())
71 feats = [n, (cnt.get("G",0)+cnt.get("C",0))/n, (cnt.get("A",0)+cnt.get("T",0))/n,
72 cnt.get("N",0)/n, max(cnt.values())/n if cnt else 0,
73 -sum((c/tot)*math.log2(c/tot) for c in cnt.values() if c>0)]
74 for k in [3,4,5,6]:
75 kc = Counter(seq[i:i+k] for i in range(n-k+1))
76 tk = max(n-k+1,1)
77 feats.extend(kc.get(km,0)/tk for km in VOCAB[k])
78 feats.extend(_codon_features(seq))
79 return feats # 5,533 features
80
81# ── Inference ─────────────────────────────────────────────────────────────────
82def screen_dna(seq, threshold_review=0.3, threshold_escalate=0.6):
83 feats = np.array([extract_features(seq)])
84 model = short_model if len(seq) < 150 else general_model
85 prob = float(model.predict_proba(feats)[0, 1])
86 if prob >= threshold_escalate: return "ESCALATE", prob
87 if prob >= threshold_review: return "REVIEW", prob
88 return "SAFE", prob
89
90decision, score = screen_dna("ATGGCTAGCATGACTGGTGGACAGCAAATGGG")
91print(f"{decision} (score: {score:.3f})")