Download Computational/package/sample_direct.py from ChatterjeeLab/Metalorian: direct link, hf CLI and curl.
- Browser
- Download file 7.61 kB
-
https://huggingface.co/ChatterjeeLab/Metalorian/resolve/main/Computational/package/sample_direct.py
- Command line
-
hf download hf://ChatterjeeLab/Metalorian/Computational/package/sample_direct.py
-
curl -L -o sample_direct.py https://huggingface.co/ChatterjeeLab/Metalorian/resolve/main/Computational/package/sample_direct.py
7.61 kB
| """Direct Metalorian sampling.""" | |
| from metalorian.sampling import Sampler, FLAGS, METAL_LABELS, TOKEN_DROPOUT_SCALE | |
| from metalorian.cli import parser, collect | |
| import torch | |
| import torch.nn.functional as F | |
| CANVAS = 1024 | |
| class FixedGradSampler(Sampler): | |
| def __init__(self, cfg, device='cuda:0'): | |
| super().__init__(cfg, device) | |
| assert float(self.aa_bias.abs().max()) == 0.0, "aa_bias must be zero (no chemistry)" | |
| print(f"guide_on={cfg.guide_on} accumulate={not cfg.no_accum} " | |
| f"grad_steps={cfg.grad_steps} scale={cfg.guidance_scale}") | |
| def x0_hat(self, x_t, t, cond): | |
| eps = self.net_sampler.model(x_t, t, cond) | |
| return self.net_sampler.predict_xstart_from_eps(x_t, t, eps=eps) | |
| def target_logit(self, lat, attn_len): | |
| cfg = self.cfg | |
| logits = self.model_con.decode_embeddings(lat[:, :attn_len]) | |
| seg = logits + self.vocab_mask[None, None, :] | |
| soft = F.softmax(seg / cfg.tau, dim=-1) | |
| emb = soft @ self.W | |
| B = emb.shape[0] | |
| cls_e = self.W[self.cls_id].view(1, 1, -1).expand(B, 1, -1) | |
| eos_e = self.W[self.eos_id].view(1, 1, -1).expand(B, 1, -1) | |
| inp = torch.cat([cls_e, emb, eos_e], dim=1) * TOKEN_DROPOUT_SCALE | |
| am = torch.ones(B, attn_len + 2, device=self.device) | |
| x_pred, _, _, _ = self.predictor.forward_from_embeds(inp, am) | |
| tgt = x_pred[:, cfg.target_idx] | |
| oth = x_pred[:, self.others].max(dim=1).values | |
| return tgt - cfg.penalty * oth, soft | |
| def init_latent(self, B, lmask): | |
| return torch.randn(B, CANVAS, 1280, device=self.device) * lmask | |
| def _encode_perpos(self, seqs): | |
| """Encode sequences as residue latents of shape (N, S, 1280).""" | |
| S = self.cfg.seq_len | |
| m = self.model_con.base_model | |
| enc = self.tokenizer(seqs, return_tensors='pt', padding='max_length', | |
| max_length=S + 2, truncation=True) | |
| ids = enc['input_ids'].to(self.device); am = enc['attention_mask'].to(self.device) | |
| emb = m.esm_model(input_ids=ids, attention_mask=am).last_hidden_state | |
| x_pool, x_attns = m.attn_head(emb, am) | |
| x_pool = m.attn_ln(x_pool + m.attn_skip(x_pool)) | |
| for lin in m.linear_layers: | |
| r = x_pool; x_pool = F.silu(lin(x_pool)); x_pool = x_pool + r | |
| x_w = torch.einsum('bhlk,bld->bhld', x_attns, x_pool).mean(dim=1) | |
| return m.clf_ln(x_w)[:, 1:1 + S, :] | |
| def _best_decode_one(self, lat_S, K, temp): | |
| """Select among K decodes by mean positional cosine similarity.""" | |
| S = self.cfg.seq_len | |
| lg = self.model_con.decode_embeddings(lat_S.unsqueeze(0))[0] + self.vocab_mask[None, :] | |
| probs = F.softmax(lg / temp, dim=-1) | |
| toks = torch.multinomial(probs, K, replacement=True).T | |
| cand = [self.tokenizer.decode(t, skip_special_tokens=True).replace(' ', '') for t in toks] | |
| cand = [c for c in cand if len(c) == S] | |
| if not cand: | |
| return self.tokenizer.decode(lg.argmax(-1), skip_special_tokens=True).replace(' ', '') | |
| xc = F.normalize(self._encode_perpos(cand), dim=-1) | |
| Ln = F.normalize(lat_S, dim=-1) | |
| cos = (xc * Ln[None]).sum(-1).mean(-1) | |
| return cand[int(cos.argmax())] | |
| def decode(self, x, attn_len): | |
| if getattr(self.cfg, 'best_decode', 0) and self.cfg.best_decode > 0: | |
| return [self._best_decode_one(x[i, :attn_len], self.cfg.best_decode, | |
| self.cfg.decode_temp) for i in range(x.shape[0])] | |
| return self.hard_decode(x, attn_len) | |
| def run_round(self): | |
| cfg, dev = self.cfg, self.device | |
| B, S = cfg.batch, cfg.seq_len | |
| lens = torch.full((B,), S, device=dev, dtype=torch.long) | |
| attn = torch.arange(CANVAS, device=dev)[None, :] < lens[:, None] | |
| lmask = attn.unsqueeze(-1).float() | |
| x = self.init_latent(B, lmask) | |
| tgt1h = torch.zeros(B, 15, device=dev) | |
| tgt1h[:, cfg.target_idx] = 1.0 | |
| log_dis = torch.log(tgt1h + 1e-10) | |
| best_score = [-1.0] * B | |
| best_seq = [None] * B | |
| cys_trace = [] | |
| for step in reversed(range(FLAGS.T)): | |
| t = torch.full((B,), step, device=dev, dtype=torch.long) | |
| for k in range(cfg.grad_steps): | |
| x = x.detach().requires_grad_(True) | |
| if cfg.guide_on == 'x0': | |
| lat = self.x0_hat(x, t, log_dis) | |
| else: | |
| lat = x | |
| obj, soft = self.target_logit(lat, S) | |
| g, = torch.autograd.grad(obj.sum(), x) | |
| with torch.no_grad(): | |
| gn = g.flatten(1).norm(dim=1).clamp_min(1e-8).view(-1, 1, 1) | |
| stepv = cfg.guidance_scale * (g / gn) * lmask | |
| if cfg.no_accum: | |
| mean, _ = self.net_sampler.p_mean_variance( | |
| x_t=x.detach(), t=t, cond=log_dis, trans=None) | |
| x = (mean * lmask + stepv).clamp(-1, 1) | |
| else: | |
| x = (x.detach() + stepv).clamp(-1, 1) | |
| x = x * lmask | |
| with torch.no_grad(): | |
| mean, log_var = self.net_sampler.p_mean_variance( | |
| x_t=x, t=t, cond=log_dis, trans=None) | |
| mean = (mean * lmask).float() | |
| noise = torch.randn_like(x) if step > 0 else torch.zeros_like(x) | |
| x = (mean + torch.exp(0.5 * log_var.float()) * noise * lmask).clamp(-1, 1) | |
| x = x * lmask | |
| log_dis = self.trainer_dis.p_sample(log_dis, t, x) | |
| if step % 2 == 0: | |
| c = cfg.label_clamp | |
| log_dis = (1 - c) * log_dis + c * torch.log(tgt1h + 1e-10) | |
| if cfg.track_best: | |
| with torch.no_grad(): | |
| cur = self.hard_decode(x, S) | |
| ok = [i for i, q in enumerate(cur) if len(q) == S] | |
| if ok: | |
| pr, _ = self.score([cur[i] for i in ok]) | |
| sc = pr[:, cfg.target_idx] | |
| for j, i in enumerate(ok): | |
| if sc[j].item() > best_score[i]: | |
| best_score[i] = sc[j].item() | |
| best_seq[i] = cur[i] | |
| if step % 10 == 0: | |
| with torch.no_grad(): | |
| seqs = self.hard_decode(x, S) | |
| cys_trace.append(sum(q.count('C') for q in seqs) / max(len(seqs), 1)) | |
| if cfg.track_best: | |
| final = self.decode(x, S) | |
| out = [best_seq[i] if best_seq[i] is not None else final[i] for i in range(B)] | |
| return out, cys_trace | |
| return self.decode(x, S), cys_trace | |
| def accept(self, seqs, _unused): | |
| cfg = self.cfg | |
| probs, fires = self.score(seqs) | |
| out = [] | |
| for i, s in enumerate(seqs): | |
| if len(s) != cfg.seq_len: | |
| continue | |
| sc = probs[i, cfg.target_idx].item() | |
| if cfg.accept_all: | |
| out.append({'sequence': s, 'score': sc, 'n_cys': s.count('C'), | |
| 'n_his': s.count('H')}) | |
| continue | |
| if sc >= self.threshold and fires[i, cfg.target_idx].item() >= 0.5: | |
| out.append({'sequence': s, 'score': sc, 'n_cys': s.count('C'), | |
| 'n_his': s.count('H')}) | |
| return out, probs | |
| if __name__ == '__main__': | |
| collect(parser('direct').parse_args(), FixedGradSampler, 'direct') | |