"""PasswordLLaMA — BBPE version. Vocab=8192, max_seq=48. Same architecture as V4: 12L-512E-8H, RoPE+SwiGLU+RMSNorm.""" import torch, torch.nn as nn, torch.nn.functional as F import math, os, json, re from safetensors.torch import load_file class RMSNorm(nn.Module): def __init__(self, dim, eps=1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x): rms = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return x * rms * self.weight def precompute_rope_freqs(dim, max_len=64, theta=10000.0): freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) t = torch.arange(max_len) freqs = torch.outer(t, freqs) cos = torch.cos(freqs).repeat_interleave(2, dim=-1) sin = torch.sin(freqs).repeat_interleave(2, dim=-1) return cos, sin def apply_rope(x, cos, sin): T = x.shape[2] cos = cos[:T].unsqueeze(0).unsqueeze(0) sin = sin[:T].unsqueeze(0).unsqueeze(0) x_real = x.float() x_rot = torch.stack([-x_real[..., 1::2], x_real[..., ::2]], dim=-1).flatten(-2) return (x_real * cos + x_rot * sin).to(x.dtype) class CausalSelfAttn(nn.Module): def __init__(self, n_embd, n_head, max_seq_len=64): super().__init__() self.n_head = n_head self.n_embd = n_embd self.head_dim = n_embd // n_head self.q_proj = nn.Linear(n_embd, n_embd, bias=False) self.k_proj = nn.Linear(n_embd, n_embd, bias=False) self.v_proj = nn.Linear(n_embd, n_embd, bias=False) self.o_proj = nn.Linear(n_embd, n_embd, bias=False) cos, sin = precompute_rope_freqs(self.head_dim, max_seq_len) self.register_buffer('rope_cos', cos, persistent=False) self.register_buffer('rope_sin', sin, persistent=False) def forward(self, x): B, T, C = x.shape q = self.q_proj(x).view(B, T, self.n_head, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(B, T, self.n_head, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(B, T, self.n_head, self.head_dim).transpose(1, 2) q = apply_rope(q, self.rope_cos, self.rope_sin) k = apply_rope(k, self.rope_cos, self.rope_sin) y = F.scaled_dot_product_attention(q, k, v, is_causal=True) y = y.transpose(1, 2).contiguous().view(B, T, C) return self.o_proj(y) class SwiGLU(nn.Module): def __init__(self, n_embd, hidden_mult=8/3): super().__init__() hidden = int(n_embd * hidden_mult) self.w1 = nn.Linear(n_embd, hidden, bias=False) self.w2 = nn.Linear(hidden, n_embd, bias=False) self.w3 = nn.Linear(n_embd, hidden, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) class Block(nn.Module): def __init__(self, n_embd, n_head, max_seq_len=64): super().__init__() self.ln1 = RMSNorm(n_embd) self.attn = CausalSelfAttn(n_embd, n_head, max_seq_len) self.ln2 = RMSNorm(n_embd) self.mlp = SwiGLU(n_embd) def forward(self, x): x = x + self.attn(self.ln1(x)) x = x + self.mlp(self.ln2(x)) return x class PasswordLLaMA(nn.Module): def __init__(self, vocab_size=8192, n_layer=12, n_embd=512, n_head=8, max_seq_len=48): super().__init__() self.wte = nn.Embedding(vocab_size, n_embd) self.blocks = nn.ModuleList([ Block(n_embd, n_head, max_seq_len) for _ in range(n_layer) ]) self.ln_f = RMSNorm(n_embd) self.lm_head = nn.Linear(n_embd, vocab_size, bias=False) self.max_seq_len = max_seq_len self.register_buffer('note', torch.zeros(0, dtype=torch.uint8), persistent=True) def load_state_dict(self, state_dict, strict=True, assign=False): state_dict = dict(state_dict) note_data = state_dict.pop('note', None) result = super().load_state_dict(state_dict, strict=False, assign=assign) if note_data is not None: self.note = note_data return result def forward(self, x, last_only=False): x = self.wte(x) for block in self.blocks: x = block(x) x = self.ln_f(x) if last_only: x = x[:, -1:] return self.lm_head(x) @torch.no_grad() def generate(self, tokenizer, temperature=1.0, top_k=50, top_p=0.0, repetition_penalty=1.0, max_len=48, min_len=4, prefix_ids=None, device='cuda'): bos_id = tokenizer.token_to_id('') eos_id = tokenizer.token_to_id('') pad_id = tokenizer.token_to_id('') self.eval() if prefix_ids is not None: ids = torch.tensor([prefix_ids], dtype=torch.long, device=device) else: ids = torch.tensor([[bos_id]], dtype=torch.long, device=device) for _ in range(max_len - ids.shape[1]): logits = self(ids)[0, -1, :] / temperature logits[pad_id] = float('-inf') if repetition_penalty != 1.0: for tok_id in ids[0]: if logits[tok_id] < 0: logits[tok_id] *= repetition_penalty else: logits[tok_id] /= repetition_penalty if top_k > 0: top_k_vals, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < top_k_vals[-1]] = float('-inf') if top_p > 0.0: sorted_logits, sorted_indices = torch.sort(logits, descending=True) cum_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) sorted_indices_to_remove = cum_probs > top_p sorted_indices_to_remove[0] = False indices_to_remove = sorted_indices[sorted_indices_to_remove] logits[indices_to_remove] = float('-inf') probs = F.softmax(logits, dim=-1) next_id = torch.multinomial(probs, 1).item() if next_id == eos_id: break ids = torch.cat([ids, torch.tensor([[next_id]], device=device)], dim=1) if ids.shape[1] >= max_len: break pw = tokenizer.decode(ids[0].tolist()) pw = pw.replace('', '').replace('', '').replace('', '').replace('', '').strip() if prefix_ids is not None: pw = pw.replace(tokenizer.decode(prefix_ids), '').strip() pw = re.sub(r'\[[A-Z]+:[^\]]*\]', '', pw) pw = re.sub(r'^.*?\]:', '', pw) # strip ][ / ]: artifacts pw = pw.lstrip(':').strip() if len(pw) < min_len: return None return pw @torch.no_grad() def generate_fast(self, tokenizer, temperature=1.0, top_k=50, max_len=48, min_len=4, prefix_ids=None, batch_size=128, device='cuda'): bos_id = tokenizer.token_to_id('') eos_id = tokenizer.token_to_id('') pad_id = tokenizer.token_to_id('') self.eval() if prefix_ids is not None: prompt = torch.tensor([prefix_ids], dtype=torch.long, device=device).repeat(batch_size, 1) else: prompt = torch.full((batch_size, 1), bos_id, dtype=torch.long, device=device) prompt_len = prompt.shape[1] max_new = max_len - prompt_len if max_new <= 0: max_new = 1 total_len = prompt_len + max_new ids = torch.full((batch_size, total_len), pad_id, dtype=torch.long, device=device) ids[:, :prompt_len] = prompt cur_len = prompt_len finished = torch.zeros(batch_size, dtype=torch.bool, device=device) for _ in range(max_new): if finished.all(): break logits = self(ids[:, :cur_len]) nxt = logits[torch.arange(batch_size), -1, :] / temperature nxt[:, pad_id] = float('-inf') if top_k > 0: vals, _ = torch.topk(nxt, min(top_k, nxt.size(-1))) nxt[nxt < vals[:, -1:]] = float('-inf') probs = F.softmax(nxt, dim=-1) nids = torch.multinomial(probs, 1).squeeze(-1) finished |= (nids == eos_id) nids[finished] = pad_id ids[:, cur_len] = nids cur_len += 1 results = [] for i in range(batch_size): pw = tokenizer.decode(ids[i].tolist()) pw = pw.replace('', '').replace('', '').replace('', '').replace('', '').strip() if prefix_ids is not None: pw = pw.replace(tokenizer.decode(prefix_ids), '').strip() pw = re.sub(r'\[[A-Z]+:[^\]]*\]', '', pw) pw = re.sub(r'^.*?\]:', '', pw) pw = pw.lstrip(':').strip() if len(pw) >= min_len: results.append(pw) return results @torch.no_grad() def generate_beam(self, tokenizer, temperature=0.8, top_k=50, max_len=48, min_len=4, prefix_ids=None, n_passwords=1000, beam=50, length_penalty=0.6, target_len=None, device='cuda'): eos_id = tokenizer.token_to_id('') pad_id = tokenizer.token_to_id('') vocab_size = tokenizer.get_vocab_size() self.eval() if prefix_ids is not None: prompt = torch.tensor([prefix_ids], dtype=torch.long, device=device) else: prompt = torch.tensor([[tokenizer.token_to_id('')]], dtype=torch.long, device=device) # Extract LEN constraint from prefix if target_len not explicitly given if target_len is None and prefix_ids is not None: prefix_str = tokenizer.decode(prefix_ids) m = re.search(r'\[LEN:(\d+)\]', prefix_str) if m: target_len = int(m.group(1)) prompt_len = prompt.shape[1] max_new = max_len - prompt_len if max_new <= 0: return [] beam_tokens = [prompt.squeeze(0).tolist()] beam_probs = [0.0] completed = [] seen_strs = set() def score(lp, l): return lp / ((max(l - prompt_len, 1)) ** length_penalty) def get_pw_len(tok_ids): """Decode password portion and return character count.""" raw = tokenizer.decode(tok_ids) raw = re.sub(r'\[[A-Z]+:[^\]]*\]', '', raw) raw = raw.replace('','').replace('','').replace('','').replace('','') raw = re.sub(r'^.*?\]:', '', raw).lstrip(':').strip() return len(raw), raw for step in range(max_new): if len(completed) >= n_passwords: break active = [(t, p) for t, p in zip(beam_tokens, beam_probs) if len(t) < prompt_len + max_new] if not active: break max_al = max(len(t[0]) for t in active) padded = torch.full((len(active), max_al), pad_id, dtype=torch.long, device=device) for i, (t, _) in enumerate(active): padded[i, :len(t)] = torch.tensor(t, device=device) logits = self(padded)[:, -1, :] new_cands = [] for i, (tok, lp) in enumerate(active): pdist = torch.nn.functional.softmax(logits[i] / temperature, dim=-1) # If target_len is set and we're already at/past it, only allow EOS if target_len is not None: cur_len, _ = get_pw_len(tok) if cur_len >= target_len: # Only allow EOS token pdist = torch.zeros_like(pdist) pdist[eos_id] = 1.0 cp, ci = torch.topk(pdist, min(top_k, vocab_size)) for j in range(len(ci)): tid = ci[j].item() nlp = lp + math.log(max(cp[j].item(), 1e-30)) if tid == eos_id: _, cl = get_pw_len(tok + [tid]) # Respect target_len and min_len min_ok = len(cl) >= min_len tgt_ok = target_len is None or len(cl) >= target_len if min_ok and tgt_ok and cl not in seen_strs: seen_strs.add(cl) completed.append((cl, score(nlp, len(tok)))) else: # If target_len is set, skip tokens that would exceed it # (except we can't easily predict, so let through) new_cands.append((tok + [tid], nlp)) if new_cands: new_cands.sort(key=lambda x: score(x[1], len(x[0])), reverse=True) beam_tokens = [c[0] for c in new_cands[:beam]] beam_probs = [c[1] for c in new_cands[:beam]] completed.sort(key=lambda x: x[1], reverse=True) return [pw for pw, _ in completed[:n_passwords]]