"""A small randomly initialized RoPE/GQA decoder with an auxiliary action head.""" from dataclasses import dataclass, asdict import math import torch from torch import nn from torch.nn import functional as F @dataclass class Config: vocab_size:int=24000 width:int=1024 layers:int=12 heads:int=16 kv_heads:int=4 hidden:int=2816 context:int=2048 rope_theta:float=10000.0 copy_dim:int=0 def to_dict(self): return asdict(self) class RMSNorm(nn.Module): def __init__(self,width): super().__init__(); self.weight=nn.Parameter(torch.ones(width)) def forward(self,x): return F.rms_norm(x,(x.shape[-1],),self.weight,eps=1e-6) def rotary(x,cos,sin): cos=cos.to(x.dtype); sin=sin.to(x.dtype) a,b=x.chunk(2,dim=-1) return torch.cat((a*cos-b*sin,b*cos+a*sin),dim=-1) class Attention(nn.Module): def __init__(self,c): super().__init__(); self.heads=c.heads; self.kv_heads=c.kv_heads; self.dim=c.width//c.heads self.qkv=nn.Linear(c.width,(c.heads+2*c.kv_heads)*self.dim,bias=False) self.out=nn.Linear(c.width,c.width,bias=False) def forward(self,x,cos,sin,past=None,pad_mask=None,use_cache=False): b,t,_=x.shape q,k,v=self.qkv(x).split([self.heads*self.dim,self.kv_heads*self.dim,self.kv_heads*self.dim],dim=-1) q=q.view(b,t,self.heads,self.dim).transpose(1,2) k=k.view(b,t,self.kv_heads,self.dim).transpose(1,2) v=v.view(b,t,self.kv_heads,self.dim).transpose(1,2) q=rotary(q,cos,sin); k=rotary(k,cos,sin) offset=0 if past is not None: offset=past[0].shape[2] k=torch.cat((past[0],k),dim=2); v=torch.cat((past[1],v),dim=2) causal=past is None mask=None if pad_mask is not None: key_positions=torch.arange(k.shape[2],device=x.device) query_positions=torch.arange(offset,offset+t,device=x.device) mask=(key_positions[None,:]<=query_positions[:,None])[None,None,:,:] & pad_mask[:,None,None,:] causal=False elif past is not None and t>1: mask=torch.arange(k.shape[2],device=x.device)[None,:]<=torch.arange(offset,offset+t,device=x.device)[:,None] y=F.scaled_dot_product_attention(q,k,v,attn_mask=mask,is_causal=causal,enable_gqa=True) y=y.transpose(1,2).contiguous().view(b,t,-1) return self.out(y),(k,v) if use_cache else None class Block(nn.Module): def __init__(self,c): super().__init__(); self.norm1=RMSNorm(c.width); self.attn=Attention(c); self.norm2=RMSNorm(c.width) self.gate_up=nn.Linear(c.width,2*c.hidden,bias=False); self.down=nn.Linear(c.hidden,c.width,bias=False) def forward(self,x,cos,sin,past=None,pad_mask=None,use_cache=False): a,cache=self.attn(self.norm1(x),cos,sin,past,pad_mask,use_cache) x=x+a; gate,up=self.gate_up(self.norm2(x)).chunk(2,dim=-1) return x+self.down(F.silu(gate)*up),cache class TinyQuery(nn.Module): def __init__(self,c): super().__init__(); self.config=c assert c.width%c.heads==0 and c.heads%c.kv_heads==0 and (c.width//c.heads)%2==0 self.tokens=nn.Embedding(c.vocab_size,c.width) self.blocks=nn.ModuleList([Block(c) for _ in range(c.layers)]) self.norm=RMSNorm(c.width); self.action_head=nn.Linear(c.width,3,bias=False) if c.copy_dim: self.copy_query=nn.Linear(c.width,c.copy_dim,bias=False) self.copy_key=nn.Linear(c.width,c.copy_dim,bias=False) self.copy_gate=nn.Linear(c.width,1) dim=c.width//c.heads inv=1/(c.rope_theta**(torch.arange(0,dim,2,dtype=torch.float32)/dim)) angles=torch.outer(torch.arange(c.context,dtype=torch.float32),inv) self.register_buffer('rope_cos',angles.cos()[None,None,:,:],persistent=False) self.register_buffer('rope_sin',angles.sin()[None,None,:,:],persistent=False) self.apply(self._init) if c.copy_dim: nn.init.zeros_(self.copy_gate.weight); nn.init.constant_(self.copy_gate.bias,2.0) for block in self.blocks: nn.init.normal_(block.attn.out.weight,std=0.02/math.sqrt(2*c.layers)) nn.init.normal_(block.down.weight,std=0.02/math.sqrt(2*c.layers)) @staticmethod def _init(m): if isinstance(m,(nn.Linear,nn.Embedding)): nn.init.normal_(m.weight,std=0.02) def forward(self,ids,targets=None,weights=None,boundaries=None,actions=None, past=None,pad_mask=None,use_cache=False,last_only=False,prompt_weight=0.15): length=ids.shape[1]; offset=0 if past is None else past[0][0].shape[2] if offset+length>self.config.context: raise ValueError('Context limit exceeded') x=self.tokens(ids) cos=self.rope_cos[:,:,offset:offset+length,:].to(x.dtype) sin=self.rope_sin[:,:,offset:offset+length,:].to(x.dtype) caches=[] for i,block in enumerate(self.blocks): x,cache=block(x,cos,sin,None if past is None else past[i],pad_mask,use_cache) if use_cache: caches.append(cache) x=self.norm(x) action_logits=None if boundaries is not None: selected=x[torch.arange(x.shape[0],device=x.device),boundaries] action_logits=self.action_head(selected) output=x[:,-1:,:] if last_only else x logits=F.linear(output,self.tokens.weight) copy_attention=None if self.config.copy_dim: keys=self.copy_key(x); source_ids=ids if past is not None: keys=torch.cat((past[-1][0],keys),dim=1) source_ids=torch.cat((past[-1][1],ids),dim=1) query=self.copy_query(output) scores=(query@keys.transpose(-1,-2)).float()/math.sqrt(self.config.copy_dim) query_positions=torch.arange(offset+length-output.shape[1],offset+length,device=ids.device) allowed=(torch.arange(keys.shape[1],device=ids.device)[None,:]<=query_positions[:,None])[None,:,:] allowed=allowed & (source_ids[:,None,:]!=0) # Copy the supplied context/question, never recycle generated response text. allowed=allowed & ((source_ids==3).cumsum(-1)==0)[:,None,:] if pad_mask is not None: allowed=allowed & pad_mask[:,None,:] copy_attention=scores.masked_fill(~allowed,-1e9).softmax(-1)*allowed gate=self.copy_gate(output).float().sigmoid() if use_cache: caches.append((keys,source_ids)) if targets is None: if copy_attention is not None: probabilities=logits.float().softmax(-1)*gate indices=source_ids[:,None,:].expand(-1,output.shape[1],-1) probabilities=probabilities.scatter_add(-1,indices,copy_attention*(1-gate)) logits=probabilities.clamp_min(1e-30).log() return logits,caches,action_logits if copy_attention is None: losses=F.cross_entropy(logits.reshape(-1,logits.shape[-1]).float(),targets.reshape(-1), ignore_index=-100,reduction='none').view_as(targets) else: safe_targets=targets.clamp_min(0) generated=logits.float().log_softmax(-1).gather(-1,safe_targets[:,:,None]).squeeze(-1).exp() copied=(copy_attention*(source_ids[:,None,:]==safe_targets[:,:,None])).sum(-1) losses=-(gate.squeeze(-1)*generated+(1-gate.squeeze(-1))*copied).clamp_min(1e-30).log() valid=targets!=-100 response=(weights>0)&valid token_weights=torch.where(response,1.0,prompt_weight)*valid lm=(losses*token_weights).sum()/token_weights.sum().clamp_min(1) auxiliary=F.cross_entropy(action_logits.float(),actions) if actions is not None else lm*0 response_loss=(losses*response).sum()/response.sum().clamp_min(1) return lm+0.05*auxiliary,torch.stack((lm.detach(),response_loss.detach(),auxiliary.detach())) @torch.no_grad() def generate_batch(self,prompts,eos_id,pad_id=0,max_new_tokens=180): self.eval(); device=next(self.parameters()).device longest=max(map(len,prompts)) if longest+max_new_tokens>self.config.context: max_new_tokens=self.config.context-longest if max_new_tokens<=0: raise ValueError('Prompt leaves no output space') ids=torch.full((len(prompts),longest),pad_id,dtype=torch.long,device=device) mask=torch.zeros_like(ids,dtype=torch.bool) for i,p in enumerate(prompts): ids[i,-len(p):]=torch.tensor(p,device=device); mask[i,-len(p):]=True outputs=[[] for _ in prompts]; finished=torch.zeros(len(prompts),dtype=torch.bool,device=device) past=None for _ in range(max_new_tokens): logits,past,_=self(ids,past=past,pad_mask=mask,use_cache=True,last_only=True) next_ids=logits[:,-1].argmax(dim=-1) done=finished.tolist() for i,token in enumerate(next_ids.tolist()): if not done[i]: outputs[i].append(token) finished|=next_ids==eos_id if finished.all(): break ids=next_ids[:,None] mask=torch.cat((mask,torch.ones((len(prompts),1),device=device,dtype=torch.bool)),dim=1) return outputs