Safetensors
Metalorian / Computational /package /sample_direct.py
yinuozhang's picture
Add computational release from Zenodo 22960654
aa45001 verified
Raw History Blame Contribute Delete
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
@torch.no_grad()
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, :]
@torch.no_grad()
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')