Image-Text-to-Text
Safetensors
GGUF
English
llama.cpp
test-fixture
tool-calling
ocr
mtp
pruning
conversational
Instructions to use Serveurperso/small-test with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- llama.cpp
How to use Serveurperso/small-test with llama.cpp:
Install (macOS, Linux)
curl -LsSf https://llama.app/install.sh | sh # Start a local OpenAI-compatible server with a web UI: llama serve -hf Serveurperso/small-test:F16 # Run inference directly in the terminal: llama cli -hf Serveurperso/small-test:F16
Install from WinGet (Windows)
winget install llama.cpp # Start a local OpenAI-compatible server with a web UI: llama serve -hf Serveurperso/small-test:F16 # Run inference directly in the terminal: llama cli -hf Serveurperso/small-test:F16
Use pre-built binary
# Download pre-built binary from: # https://github.com/ggerganov/llama.cpp/releases # Start a local OpenAI-compatible server with a web UI: ./llama-server -hf Serveurperso/small-test:F16 # Run inference directly in the terminal: ./llama-cli -hf Serveurperso/small-test:F16
Build from source code
git clone https://github.com/ggerganov/llama.cpp.git cd llama.cpp cmake -B build cmake --build build -j --target llama-server llama-cli # Start a local OpenAI-compatible server with a web UI: ./build/bin/llama-server -hf Serveurperso/small-test:F16 # Run inference directly in the terminal: ./build/bin/llama-cli -hf Serveurperso/small-test:F16
Use Docker
docker model run hf.co/Serveurperso/small-test:F16
- LM Studio
- Jan
- vLLM
How to use Serveurperso/small-test with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "Serveurperso/small-test" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Serveurperso/small-test", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker
docker model run hf.co/Serveurperso/small-test:F16
- Ollama
How to use Serveurperso/small-test with Ollama:
ollama run hf.co/Serveurperso/small-test:F16
- Unsloth Desktop
- Pi
How to use Serveurperso/small-test with Pi:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf Serveurperso/small-test:F16
Configure the model in Pi
# Install Pi: npm install -g @earendil-works/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "llama-cpp": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "Serveurperso/small-test:F16" } ] } } }Run Pi
# Start Pi in your project directory: pi
- Docker Model Runner
How to use Serveurperso/small-test with Docker Model Runner:
docker model run hf.co/Serveurperso/small-test:F16
- Lemonade
How to use Serveurperso/small-test with Lemonade:
Pull the model
# Download Lemonade from https://lemonade-server.ai/ lemonade pull Serveurperso/small-test:F16
Run and chat with the model
lemonade run user.small-test-F16
List all available models
lemonade list
- Hermes Agent
How to use Serveurperso/small-test with Hermes Agent:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf Serveurperso/small-test:F16
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default Serveurperso/small-test:F16
Run Hermes
hermes
- Atomic Chat
- OpenClaw
How to use Serveurperso/small-test with OpenClaw:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf Serveurperso/small-test:F16
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "Serveurperso/small-test:F16" \ --custom-provider-id llama-cpp \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
Download src/train.py from Serveurperso/small-test: direct link, hf CLI and curl.
- Browser
- Download file 21.5 kB
-
https://huggingface.co/Serveurperso/small-test/resolve/main/src/train.py
- Command line
-
hf download hf://Serveurperso/small-test/src/train.py
-
curl -L -o train.py https://huggingface.co/Serveurperso/small-test/resolve/main/src/train.py
21.5 kB
| import torch, json, os, sys, time, math, argparse, random, copy, glob | |
| import numpy as np | |
| import torch.nn as nn, torch.nn.functional as F | |
| from torch.utils.tensorboard import SummaryWriter | |
| from transformers import AutoTokenizer, Qwen3_5ForConditionalGeneration | |
| from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5DecoderLayer, Qwen3_5RMSNorm | |
| from safetensors.torch import load_file, save_file | |
| from muon import Muon | |
| ap=argparse.ArgumentParser() | |
| ap.add_argument("--model",required=True) | |
| ap.add_argument("--data",default="tokdata") | |
| ap.add_argument("--mix",default="smoltalk:0.55,everyday:0.05,apigen:0.20,hermes_fc:0.04,hermes_fc_single:0.04,hermes_glaive:0.12") | |
| ap.add_argument("--out",required=True) | |
| ap.add_argument("--seq",type=int,default=2048) | |
| ap.add_argument("--bs",type=int,default=8) | |
| ap.add_argument("--accum",type=int,default=1) | |
| ap.add_argument("--steps",type=int,default=1000,help="schedule length") | |
| ap.add_argument("--stop",type=int,default=None,help="last step of this phase, the schedule continues on resume") | |
| ap.add_argument("--warmup",type=int,default=100) | |
| ap.add_argument("--lr-muon",type=float,default=2e-3) | |
| ap.add_argument("--lr-adam",type=float,default=1e-3) | |
| ap.add_argument("--lr-min-ratio",type=float,default=0.1) | |
| ap.add_argument("--wd",type=float,default=0.01) | |
| ap.add_argument("--mtp-weight",type=float,default=0.3) | |
| ap.add_argument("--eval-every",type=int,default=200) | |
| ap.add_argument("--save-every",type=int,default=1000) | |
| ap.add_argument("--gen-every",type=int,default=500) | |
| ap.add_argument("--max-minutes",type=float,default=1e9) | |
| ap.add_argument("--logdir",default="runs") | |
| ap.add_argument("--name",default=None) | |
| ap.add_argument("--seed",type=int,default=0) | |
| ap.add_argument("--resume",default=None) | |
| ap.add_argument("--profile",action="store_true") | |
| ap.add_argument("--reset-step",action="store_true",help="resume weights only, start the schedule from step 0") | |
| ap.add_argument("--ocr-frac",type=float,default=0.0,help="fraction of packed docs that are synthetic OCR image docs") | |
| ap.add_argument("--vis-lr-mult",type=float,default=0.5) | |
| ap.add_argument("--ocr-workers",type=int,default=6) | |
| ap.add_argument("--teacher",default=None,help="unpruned parent model, logit distillation on the assistant tokens") | |
| ap.add_argument("--kd-weight",type=float,default=0.5,help="share of the distillation loss in the main loss") | |
| args=ap.parse_args() | |
| torch.manual_seed(args.seed); random.seed(args.seed); np.random.seed(args.seed) | |
| dev="cuda" | |
| name=args.name or time.strftime("%m%d-%H%M") | |
| writer=SummaryWriter(os.path.join(args.logdir,name)) | |
| os.makedirs(args.out,exist_ok=True) | |
| # ---------------- data | |
| class Source: | |
| def __init__(s,path): | |
| z=np.load(path); s.ids=z["ids"]; s.mask=z["mask"]; s.offs=z["offs"]; s.n=len(s.offs)-1; s.perm=np.random.permutation(s.n); s.pos=0 | |
| def next_doc(s): | |
| if s.pos>=s.n: s.perm=np.random.permutation(s.n); s.pos=0 | |
| i=s.perm[s.pos]; s.pos+=1 | |
| a,b=s.offs[i],s.offs[i+1] | |
| return s.ids[a:b],s.mask[a:b] | |
| mix=[(k,float(v)) for k,v in (x.split(":") for x in args.mix.split(","))] | |
| srcs={k:Source(os.path.join(args.data,k+".npz")) for k,_ in mix} | |
| weights=np.array([w for k,w in mix]); weights/=weights.sum() # absolute sampling fractions (per doc) | |
| print("mix:",{k:round(float(p),3) for (k,_),p in zip(mix,weights)}) | |
| # ---- OCR image docs, produced in worker processes | |
| _tok=None; _ip=None; _tpl=None; IMG_ID=None | |
| def _ocr_init(model_dir): | |
| global _tok,_ip,_tpl,IMG_ID | |
| from transformers import AutoTokenizer, AutoImageProcessor | |
| import render as R | |
| _tok=AutoTokenizer.from_pretrained(model_dir); _ip=AutoImageProcessor.from_pretrained(model_dir) | |
| R.init(model_dir); _tpl=R | |
| IMG_ID=_tok.convert_tokens_to_ids("<|image_pad|>") | |
| def ocr_doc(seed): | |
| import random as _r | |
| from ocr_gen import make_sample, preprocess | |
| rng=_r.Random(seed); img,msgs=make_sample(rng) | |
| enc=preprocess(img,_ip,rng) | |
| grid=enc["image_grid_thw"][0]; n=int(grid[0]*grid[1]*grid[2])//4 | |
| text,spans=_tpl.render(msgs) | |
| e=_tok(text,add_special_tokens=False,return_offsets_mapping=True) | |
| ids=e["input_ids"]; offs=e["offset_mapping"] | |
| m=np.zeros(len(ids),dtype=np.uint8) | |
| st=np.array([o[0] for o in offs]); en=np.array([o[1] for o in offs]) | |
| for a,b in spans: m[(en>a)&(st<b)]=1 | |
| i=ids.index(IMG_ID) | |
| ids=np.array(ids[:i]+[IMG_ID]*n+ids[i+1:],dtype=np.int32); m=np.concatenate([m[:i],np.zeros(n,dtype=np.uint8),m[i+1:]]) | |
| return ids,m,enc["pixel_values"].astype(np.float16),grid.astype(np.int64) | |
| ocr_iter=None | |
| def _ocr_worker(q,model_dir,seed0): | |
| # bounded queue: never produce more docs than the trainer consumes | |
| _ocr_init(model_dir); i=seed0 | |
| while True: | |
| q.put(ocr_doc(i)); i+=1 | |
| if args.ocr_frac>0: | |
| import multiprocessing as mp | |
| ctx=mp.get_context("fork"); ocr_q=ctx.Queue(maxsize=64) | |
| for w in range(args.ocr_workers): | |
| p=ctx.Process(target=_ocr_worker,args=(ocr_q,args.model,args.seed*10**7+w*10**6),daemon=True); p.start() | |
| def _ocr_gen(): | |
| while True: yield ocr_q.get() | |
| ocr_iter=_ocr_gen() | |
| def pack(next_text_doc,next_ocr_doc,bs,ocr_frac,rng=np.random,pending=None): | |
| # whole documents only: a document that does not fit waits in pending for a later row, | |
| # the remainder of a row is padded with zero loss once no pending document fits | |
| L=args.seq+2; pending=[] if pending is None else pending | |
| ids=np.zeros((bs,L),dtype=np.int64); mask=np.zeros((bs,L),dtype=np.float32); pv=[]; grids=[] | |
| for b in range(bs): | |
| cur=0 | |
| while cur<L: | |
| d=None | |
| for i,(pd,pm) in enumerate(pending): | |
| if len(pd)<=L-cur: d,m=pending.pop(i); break | |
| if d is None: | |
| if ocr_frac>0 and rng.random()<ocr_frac: | |
| od=next_ocr_doc() | |
| if len(od[0])<=L-cur: d,m=od[0],od[1]; pv.append(od[2]); grids.append(od[3]) | |
| if d is None: | |
| d,m=next_text_doc() | |
| if len(d)>L: d,m=d[:L],m[:L] | |
| if len(d)>L-cur: | |
| pending.append((d,m)) | |
| if len(pending)<8: continue | |
| break | |
| n=len(d); ids[b,cur:cur+n]=d; mask[b,cur:cur+n]=m; cur+=n | |
| out=(torch.from_numpy(ids),torch.from_numpy(mask)) | |
| if pv: out=out+(torch.from_numpy(np.concatenate(pv)),torch.from_numpy(np.stack(grids))) | |
| else: out=out+(None,None) | |
| return out | |
| def next_text_doc(): | |
| k=mix[np.random.choice(len(mix),p=weights)][0]; return srcs[k].next_doc() | |
| train_pending=[] | |
| def make_batch(): | |
| return pack(next_text_doc,(lambda: next(ocr_iter)) if ocr_iter else None,args.bs,args.ocr_frac,pending=train_pending) | |
| eval_sets={} | |
| for f in sorted(glob.glob(os.path.join(args.data,"*.eval.npz")))+[os.path.join(args.data,"smoltalk_eval.npz")]: | |
| s=Source(f); k=os.path.basename(f).replace(".npz","") | |
| eval_sets[k]=[pack(s.next_doc,None,8,0.0) for _ in range(4)] | |
| if args.ocr_frac>0: | |
| _ocr_init(args.model) | |
| import itertools; docs=[ocr_doc(-1-i) for i in range(96)]; it=itertools.cycle(docs) | |
| ev=Source(os.path.join(args.data,"everyday.eval.npz")) | |
| eval_sets["ocr.eval"]=[pack(ev.next_doc,lambda: next(it),4,0.9,np.random.RandomState(2)) for _ in range(3)] | |
| # ---------------- model | |
| tok=AutoTokenizer.from_pretrained(args.model) | |
| model=Qwen3_5ForConditionalGeneration.from_pretrained(args.model,dtype=torch.float32,attn_implementation="sdpa").to(dev) | |
| lm=model.model.language_model; tcfg=lm.config | |
| class MTP(nn.Module): | |
| def __init__(s,cfg): | |
| super().__init__() | |
| c=copy.deepcopy(cfg); c.layer_types=["full_attention"]; c.num_hidden_layers=1 | |
| s.fc=nn.Linear(2*cfg.hidden_size,cfg.hidden_size,bias=False) | |
| s.pre_fc_norm_embedding=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) | |
| s.pre_fc_norm_hidden=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) | |
| s.layers=nn.ModuleList([Qwen3_5DecoderLayer(c,0)]) | |
| s.norm=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) | |
| def forward(s,h,emb_next,rotary,pos=None): | |
| # h: normed trunk output for tokens t, emb_next: embedding of token t+1 -> predicts t+2 | |
| x=s.fc(torch.cat([s.pre_fc_norm_embedding(emb_next),s.pre_fc_norm_hidden(h)],-1)) | |
| if pos is None: pos=torch.arange(x.shape[1],device=x.device).view(1,1,-1).expand(3,x.shape[0],-1) | |
| elif pos.shape[0]==4: pos=pos[1:] | |
| pe=rotary(x,pos) | |
| x=s.layers[0](x,position_embeddings=pe,attention_mask=None,position_ids=pos[0]) | |
| return s.norm(x) | |
| mtp=MTP(tcfg).to(dev) | |
| mtp_sd={k[4:]:v.float() for k,v in load_file(os.path.join(args.model,"model.safetensors")).items() if k.startswith("mtp.")} | |
| print("mtp load:",mtp.load_state_dict(mtp_sd,strict=True)) | |
| for p in model.model.visual.parameters(): p.requires_grad_(args.ocr_frac>0) | |
| step0=0 | |
| if args.resume: | |
| ck=torch.load(args.resume,map_location=dev); model.load_state_dict(ck["model"]); mtp.load_state_dict(ck["mtp"]); step0=0 if args.reset_step else ck["step"] | |
| # compile the full-attention blocks (training shapes only; eval/generate stays eager) | |
| def _compile_layer(l): | |
| eager=l.forward; comp=torch.compile(eager,dynamic=False) | |
| l.forward=lambda *a,**k: comp(*a,**k) if l.training else eager(*a,**k) | |
| for l in lm.layers: | |
| if l.block_type=="full_attention": _compile_layer(l) | |
| _compile_layer(mtp.layers[0]) | |
| # ---------------- teacher: the pruned vocabulary keeps the token strings, so its ids map onto the teacher's | |
| teacher=None | |
| if args.teacher: | |
| teacher=Qwen3_5ForConditionalGeneration.from_pretrained(args.teacher,dtype=torch.bfloat16,attn_implementation="sdpa").to(dev).eval() | |
| for p in teacher.parameters(): p.requires_grad_(False) | |
| ttok=AutoTokenizer.from_pretrained(args.teacher) | |
| n_real=len(tok) | |
| old_ids=torch.tensor(ttok.convert_tokens_to_ids(tok.convert_ids_to_tokens(list(range(n_real)))),device=dev) | |
| assert (old_ids>=0).all() | |
| old_ids=torch.cat([old_ids,old_ids[:1].expand(model.lm_head.weight.shape[0]-n_real)]) | |
| W_t=teacher.lm_head.weight[old_ids[:n_real]].contiguous() | |
| # the teacher is a thinking model and our data has no reasoning: its think tokens are removed from the target | |
| think_ids=torch.tensor(tok.convert_tokens_to_ids(["<think>","</think>"]),device=dev) | |
| print(f"teacher {sum(p.numel() for p in teacher.parameters())/1e6:.0f}M, kept vocab {n_real}") | |
| def kd_loss(h,th,m,chunk=4096): | |
| # KL(teacher || student) over the kept vocabulary on the assistant positions, one chunk of logits at a time | |
| idx=m.reshape(-1).nonzero().squeeze(1) | |
| hs=h.reshape(-1,h.shape[-1])[idx]; ht=th.reshape(-1,th.shape[-1])[idx]; W=model.lm_head.weight[:n_real] | |
| tot=0 | |
| for i in range(0,len(idx),chunk): | |
| with torch.no_grad(): | |
| t=F.linear(ht[i:i+chunk],W_t).float(); t[:,think_ids]=-1e4; t=F.log_softmax(t,-1) | |
| s=F.log_softmax(F.linear(hs[i:i+chunk],W).float(),-1) | |
| tot=tot+(t.exp()*(t-s)).sum() | |
| return tot/max(1,len(idx)) | |
| n_text=sum(p.numel() for p in lm.parameters()); n_mtp=sum(p.numel() for p in mtp.parameters()) | |
| print(f"params text {n_text/1e6:.2f}M mtp {n_mtp/1e6:.2f}M total {(n_text+n_mtp)/1e6:.2f}M") | |
| # ---------------- optimizers: Muon for 2D hidden matrices, AdamW for the rest | |
| muon_p=[]; adam_p=[]; adam_emb=[] | |
| for n,p in list(lm.named_parameters())+[("mtp."+n,p) for n,p in mtp.named_parameters()]: | |
| if not p.requires_grad: continue | |
| if "embed_tokens" in n: adam_emb.append(p) | |
| elif p.ndim==2 and "norm" not in n: muon_p.append(p) | |
| else: adam_p.append(p) | |
| print(f"muon params {sum(p.numel() for p in muon_p)/1e6:.1f}M, adam {sum(p.numel() for p in adam_p)/1e6:.2f}M, emb {sum(p.numel() for p in adam_emb)/1e6:.1f}M") | |
| vis_muon=[p for n,p in model.model.visual.named_parameters() if p.requires_grad and p.ndim==2 and "norm" not in n] | |
| vis_adam=[p for n,p in model.model.visual.named_parameters() if p.requires_grad and not (p.ndim==2 and "norm" not in n)] | |
| mgroups=[{"params":muon_p,"lr_mult":1.0}]+([{"params":vis_muon,"lr_mult":args.vis_lr_mult}] if vis_muon else []) | |
| agroups=[{"params":adam_p,"weight_decay":0.0,"lr_mult":1.0},{"params":adam_emb,"weight_decay":args.wd,"lr_mult":1.0}]+([{"params":vis_adam,"weight_decay":0.0,"lr_mult":args.vis_lr_mult}] if vis_adam else []) | |
| opt_m=Muon(mgroups,lr=args.lr_muon,momentum=0.95,weight_decay=args.wd) | |
| opt_a=torch.optim.AdamW(agroups,lr=args.lr_adam,betas=(0.9,0.95),fused=True) | |
| all_params=list(lm.parameters())+list(mtp.parameters())+vis_muon+vis_adam | |
| print(f"vision trainable {sum(p.numel() for p in vis_muon+vis_adam)/1e6:.1f}M") | |
| def lr_mult(step): | |
| if step<args.warmup: return (step+1)/args.warmup | |
| t=(step-args.warmup)/max(1,args.steps-args.warmup) | |
| return args.lr_min_ratio+(1-args.lr_min_ratio)*0.5*(1+math.cos(math.pi*min(1,t))) | |
| from torch.utils.checkpoint import checkpoint | |
| def _ce_chunk(h,W,tgt,m): | |
| logits=F.linear(h,W).float() | |
| return (F.cross_entropy(logits,tgt,reduction="none")*m).sum() | |
| from liger_kernel.transformers.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss | |
| _lce=LigerFusedLinearCrossEntropyLoss(ignore_index=-100,reduction="mean") | |
| def ce_fused(h,W,tgt,m): | |
| return _lce(W,h,torch.where(m>0,tgt,torch.full_like(tgt,-100))) | |
| def ce_chunked(h,W,tgt,m,chunk=2048): | |
| # h: [N,D], tgt/m: [N]; only one chunk of logits is materialized at a time | |
| tot=0 | |
| for i in range(0,h.shape[0],chunk): | |
| tot=tot+checkpoint(_ce_chunk,h[i:i+chunk],W,tgt[i:i+chunk],m[i:i+chunk],use_reentrant=False) | |
| return tot/m.sum().clamp(min=1) | |
| IMG_ID=tok.convert_tokens_to_ids("<|image_pad|>") | |
| def compute_loss(ids,mask,pv=None,grid=None,kd=True): | |
| ids=ids.to(dev); mask=mask.to(dev) | |
| x=ids[:,:-2] | |
| with torch.autocast("cuda",dtype=torch.bfloat16): | |
| if pv is not None: | |
| mm=(x==IMG_ID).long(); pv=pv.to(dev); grid=grid.to(dev) | |
| pos=model.model.compute_3d_position_ids(input_ids=x,inputs_embeds=None,image_grid_thw=grid,mm_token_type_ids=mm) | |
| out=model.model(input_ids=x,pixel_values=pv,image_grid_thw=grid,mm_token_type_ids=mm,position_ids=pos,use_cache=False) | |
| else: | |
| pos=None | |
| out=lm(input_ids=x,use_cache=False) | |
| h=out.last_hidden_state | |
| W=model.lm_head.weight | |
| loss1=ce_fused(h.reshape(-1,h.shape[-1]),W,ids[:,1:-1].reshape(-1),mask[:,1:-1].reshape(-1)) | |
| if teacher is not None and kd: | |
| with torch.no_grad(): | |
| tx=old_ids[x] | |
| if pv is not None: th=teacher.model(input_ids=tx,pixel_values=pv,image_grid_thw=grid,mm_token_type_ids=mm,position_ids=pos,use_cache=False).last_hidden_state | |
| else: th=teacher.model.language_model(input_ids=tx,use_cache=False).last_hidden_state | |
| loss1=(1-args.kd_weight)*loss1+args.kd_weight*kd_loss(h,th,mask[:,1:-1]) | |
| emb_next=lm.embed_tokens(ids[:,1:-1]) | |
| h2=mtp(h,emb_next,lm.rotary_emb,pos) | |
| loss2=ce_fused(h2.reshape(-1,h2.shape[-1]),W,ids[:,2:].reshape(-1),mask[:,2:].reshape(-1)) | |
| return loss1,loss2 | |
| def evaluate(step): | |
| model.eval(); mtp.eval(); res={} | |
| for k,batches in eval_sets.items(): | |
| a=b=0 | |
| for ids,mask,pv,grid in batches: | |
| l1,l2=compute_loss(ids,mask,pv,grid,kd=False); a+=l1.item(); b+=l2.item() | |
| res[k]=(a/len(batches),b/len(batches)) | |
| writer.add_scalar(f"eval/{k}",a/len(batches),step); writer.add_scalar(f"eval_mtp/{k}",b/len(batches),step) | |
| model.train(); mtp.train() | |
| print(f"[eval step {step}] "+" ".join(f"{k}={v[0]:.3f}/{v[1]:.3f}" for k,v in res.items()),flush=True) | |
| GEN_PROMPTS=[ | |
| [{"role":"user","content":"Hello! Who are you?"}], | |
| [{"role":"user","content":"Write a short poem about the sea."}], | |
| [{"role":"user","content":"What is the capital of France?"}], | |
| ] | |
| GEN_TOOLS=[{"type":"function","function":{"name":"get_weather","description":"Get the current weather for a city","parameters":{"type":"object","properties":{"city":{"type":"string","description":"City name"}},"required":["city"]}}}] | |
| GEN_TOOL_PROMPTS=[ | |
| ([{"role":"user","content":"What's the weather like in Paris right now?"}],GEN_TOOLS), | |
| ([{"role":"user","content":"What's the weather in Tokyo?"},{"role":"assistant","content":"","tool_calls":[{"type":"function","function":{"name":"get_weather","arguments":{"city":"Tokyo"}}}]},{"role":"tool","content":"{\"temperature\": 22, \"condition\": \"sunny\"}"}],GEN_TOOLS), | |
| ] | |
| def generate_samples(step): | |
| model.eval(); text="" | |
| for msgs,tools in [(m,None) for m in GEN_PROMPTS]+GEN_TOOL_PROMPTS: | |
| prompt=tok.apply_chat_template(msgs,tools=tools,tokenize=False,add_generation_prompt=True) | |
| x=tok(prompt,return_tensors="pt",add_special_tokens=False).input_ids.to(dev) | |
| with torch.autocast("cuda",dtype=torch.bfloat16): | |
| y=model.generate(input_ids=x,max_new_tokens=120,do_sample=False,eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"),pad_token_id=tok.convert_tokens_to_ids("<|endoftext|>")) | |
| out=tok.decode(y[0,x.shape[1]:]) | |
| text+=f"### {msgs[-1]['role']}: {str(msgs[-1]['content'])[:80]}\n\n```\n{out}\n```\n\n" | |
| if args.ocr_frac>0: | |
| import random as _r; from ocr_gen import make_sample, preprocess | |
| from transformers import AutoImageProcessor | |
| ip=AutoImageProcessor.from_pretrained(args.model) | |
| for sd in [7,8]: | |
| img,msgs=make_sample(_r.Random(sd)); enc=preprocess(img,ip,return_tensors="pt"); g=enc["image_grid_thw"][0]; n=int(g[0]*g[1]*g[2])//4 | |
| prompt=tok.apply_chat_template(msgs[:1],tokenize=False,add_generation_prompt=True) | |
| ids=tok(prompt,add_special_tokens=False)["input_ids"]; i=ids.index(IMG_ID); ids=ids[:i]+[IMG_ID]*n+ids[i+1:] | |
| x=torch.tensor([ids],device=dev); mm=(x==IMG_ID).long() | |
| with torch.autocast("cuda",dtype=torch.bfloat16): | |
| y=model.generate(input_ids=x,pixel_values=enc["pixel_values"].to(dev),image_grid_thw=enc["image_grid_thw"].to(dev),mm_token_type_ids=mm,max_new_tokens=60,do_sample=False,eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"),pad_token_id=tok.convert_tokens_to_ids("<|endoftext|>")) | |
| out=tok.decode(y[0,x.shape[1]:]) | |
| text+=f"### OCR truth: {msgs[1]['content'][:80]!r}\n\n```\n{out}\n```\n\n" | |
| writer.add_text("samples",text,step); model.train() | |
| print(text,flush=True) | |
| def save(step,final=False): | |
| sd={k:v.detach().to(torch.bfloat16).cpu().contiguous() for k,v in model.state_dict().items() if not k.startswith("lm_head.")} | |
| sd.update({"mtp."+k:v.detach().to(torch.bfloat16).cpu().contiguous() for k,v in mtp.state_dict().items()}) | |
| d=args.out if final else os.path.join(args.out,f"step{step}") | |
| os.makedirs(d,exist_ok=True) | |
| save_file(sd,os.path.join(d,"model.safetensors"),metadata={"format":"pt"}) | |
| for f in glob.glob(os.path.join(args.model,"*.json"))+glob.glob(os.path.join(args.model,"*.jinja")): | |
| if "index" not in f: os.system(f"cp {f} {d}/") | |
| torch.save({"model":model.state_dict(),"mtp":mtp.state_dict(),"step":step},os.path.join(args.out,"resume.pt")) | |
| print("saved",d,flush=True) | |
| # ---------------- loop | |
| model.train(); mtp.train() | |
| t0=time.time(); tokens=0; tstart=time.time() | |
| evaluate(step0) | |
| if args.profile: | |
| import collections; T=collections.defaultdict(float) | |
| def sync(): torch.cuda.synchronize(); return time.time() | |
| for it in range(8): | |
| t=sync(); ids,mask,pv,grid=make_batch(); t1=sync(); T["data"]+=t1-t | |
| loss1,loss2=compute_loss(ids,mask,pv,grid); t2=sync(); T["fwd"]+=t2-t1 | |
| (loss1+args.mtp_weight*loss2).backward(); t3=sync(); T["bwd"]+=t3-t2 | |
| gn=torch.nn.utils.clip_grad_norm_(all_params,1.0); t4=sync(); T["clip"]+=t4-t3 | |
| opt_m.step(); t5=sync(); T["muon"]+=t5-t4 | |
| opt_a.step(); opt_m.zero_grad(set_to_none=True); opt_a.zero_grad(set_to_none=True); t6=sync(); T["adam"]+=t6-t5 | |
| if it==1: T.clear() | |
| print({k:round(v/6*1000,1) for k,v in T.items()},"ms/step; tokens/step",ids.numel()); sys.exit() | |
| for step in range(step0,args.steps): | |
| m=lr_mult(step) | |
| for g in opt_m.param_groups: g["lr"]=args.lr_muon*m*g["lr_mult"] | |
| for g in opt_a.param_groups: g["lr"]=args.lr_adam*m*g["lr_mult"] | |
| for _ in range(args.accum): | |
| ids,mask,pv,grid=make_batch() | |
| loss1,loss2=compute_loss(ids,mask,pv,grid) | |
| loss=(loss1+args.mtp_weight*loss2)/args.accum | |
| loss.backward() | |
| tokens+=ids.numel() | |
| gn=torch.nn.utils.clip_grad_norm_(all_params,1.0) | |
| opt_m.step(); opt_a.step(); opt_m.zero_grad(set_to_none=True); opt_a.zero_grad(set_to_none=True) | |
| if step%10==0: | |
| dt=time.time()-t0; tps=tokens/dt; t0=time.time(); tokens=0 | |
| writer.add_scalar("train/loss",loss1.item(),step); writer.add_scalar("train/loss_mtp",loss2.item(),step) | |
| writer.add_scalar("train/lr_muon",args.lr_muon*m,step); writer.add_scalar("train/grad_norm",gn.item(),step); writer.add_scalar("train/tokens_per_s",tps,step) | |
| import resource; rss=resource.getrusage(resource.RUSAGE_SELF).ru_maxrss/1e6 | |
| print(f"step {step} loss {loss1.item():.4f} mtp {loss2.item():.4f} gn {gn.item():.2f} lr {args.lr_muon*m:.2e} tok/s {tps:.0f} elapsed {(time.time()-tstart)/60:.1f}m rss {rss:.1f}G",flush=True) | |
| if (step+1)%args.eval_every==0: evaluate(step+1) | |
| if (step+1)%args.gen_every==0: generate_samples(step+1) | |
| if (step+1)%args.save_every==0 and step+1<args.steps: save(step+1) | |
| if (time.time()-tstart)/60>args.max_minutes: print("time limit reached"); break | |
| if args.stop and step+1>=args.stop: break | |
| evaluate(step+1); generate_samples(step+1); save(step+1,final=True) | |
| writer.close() | |