import json, math from pathlib import Path import torch from .model import EXLLM,EXLLMConfig from .tokenizer import HybridTokenizer,UTF8State,normalize_text def load_model(ckpt='weights/EXLLM-v1.1-5m-release3.pt'): tok=HybridTokenizer.load('tokenizer.json'); c=torch.load(ckpt,map_location='cpu'); cfg=EXLLMConfig(**c['config']); m=EXLLM(cfg); m.load_state_dict(c['model']); m.eval(); return m,tok,c def bad_text(s): if not s or '\ufffd' in s or '\x00' in s: return True if any((ord(c)<32 and c not in '\n\t') for c in s): return True # obvious repetition collapse if len(s)>=12: for n in range(1,7): unit=s[-n:] if len(unit)*4<=len(s) and s.endswith(unit*4): return True return False @torch.inference_mode() def generate(m,tok,prompt,max_new=64,temperature=0.0,top_k=12,confidence_fallback=True): ids=tok.encode_user(prompt) # Reserve output space. Left-trim user content but preserve role tokens. reserve=min(max_new, m.cfg.max_seq_len//2) if len(ids)>m.cfg.max_seq_len-reserve: keep=m.cfg.max_seq_len-reserve-3 body=ids[2:-1][-max(1,keep):]; ids=[tok.BOS,tok.USER,*body,tok.ASSIST] out=[]; state=UTF8State(); lps=[] for _ in range(min(max_new,m.cfg.max_seq_len-len(ids))): x=torch.tensor([ids+out],dtype=torch.long); logits=m(x)[0,-1].clone() # Structural specials are never generated inside assistant text. logits[tok.PAD]=logits[tok.BOS]=logits[tok.USER]=logits[tok.ASSIST]=-1e30 # UTF-8 constrained token mask. valid=torch.zeros_like(logits,dtype=torch.bool) if state.complete: valid[tok.EOS]=True for t in range(256): if state.accepts_byte(t): valid[t]=True if tok.chars: valid[256:256+len(tok.chars)]=True else: for t in range(256): if state.accepts_byte(t): valid[t]=True logits[~valid]=-1e30 lp=torch.log_softmax(logits,dim=-1) if temperature and temperature>0: z=logits/temperature if top_k and top_k