"""Bounded upstream-MLX research harness, not a production training service. Cosmos-Reason2 local conversion; frozen visual features; explicit two-stage forward/backward over TCP. Removal condition: qualified fleet training runtime. No pickle or remote executable payloads. No trained-model quality claim. """ import argparse, hashlib, io, json, os, resource, socket, struct, time from pathlib import Path import numpy as np import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim from mlx.utils import tree_flatten, tree_unflatten from mlx_vlm import load from mlx_vlm.trainer.lora_layers import LoRALinear def emit(event, **kw): print(json.dumps(dict(event=event, **kw)), flush=True) def flat(tree): return dict(tree_flatten(tree)) def save_tree(path, tree): mx.save_safetensors(str(path), flat(tree)) def load_tree(path): return tree_unflatten(list(mx.load(str(path)).items())) def exact(s, n): result = bytearray() while len(result) < n: b = s.recv(n - len(result)) if not b: raise EOFError('peer disconnected before message completed') result.extend(b) return result def send(s, meta, tensors=None): b = io.BytesIO() np.savez(b, **{k:np.array(v) for k,v in (tensors or {}).items()}) header = json.dumps(meta).encode(); payload = b.getvalue() s.sendall(struct.pack('!QQ', len(header), len(payload)) + header + payload) return 16 + len(header) + len(payload) def recv(s): a,b = struct.unpack('!QQ', exact(s,16)) if a > 65536 or b > 128*1024*1024: raise ValueError('message exceeds bounds') meta = json.loads(exact(s,a)) with np.load(io.BytesIO(exact(s,b)),allow_pickle=False) as z: data = {k:mx.array(z[k]) for k in z.files} return meta,data class Stage(nn.Module): def __init__(self, lm, start, end, last=False): super().__init__() self.layers = lm.model.layers[start:end] self.start = start if last: self.norm = lm.model.norm self.head = lm.lm_head def __call__(self, h, batch): for offset, layer in enumerate(self.layers): h = layer(h, mask='causal', position_ids=batch['position_ids']) idx = self.start + offset if idx < 3 and 'deepstack' in batch: mask = batch['visual_mask'] # Data preparation emits dense additions, avoiding host indexing # inside the differentiated path. h = h + batch['deepstack'][idx].astype(h.dtype) if 'head' in self: h = self.head(self.norm(h)) return h def loss(logits, batch): ce = nn.losses.cross_entropy(logits.astype(mx.float32), batch['labels'], reduction='none') return (ce * batch['loss_mask']).sum() / batch['loss_mask'].sum() def make_stages(base): mx.random.seed(20260909) model, processor = load(base) model.freeze() for layer in model.language_model.model.layers: for key in ('q_proj','v_proj'): original = getattr(layer.self_attn,key) setattr(layer.self_attn,key,LoRALinear.from_base(original,r=4,dropout=0,scale=2.0)) first=Stage(model.language_model,0,14) last=Stage(model.language_model,14,28,last=True) mx.eval(first.parameters(),last.parameters()) return model,processor,first,last def prepare(args): from mlx_vlm.prompt_utils import apply_chat_template from mlx_vlm.utils import prepare_inputs model, processor=load(args.base) row=json.loads(Path(args.dataset).read_text().splitlines()[0]) full=apply_chat_template(processor,model.config,row['messages'],num_images=1,add_generation_prompt=False) prefix=apply_chat_template(processor,model.config,row['messages'][:-1],num_images=1,add_generation_prompt=True) kwargs=dict(processor=processor,images=row['images'],image_token_index=model.config.image_token_index) inp=prepare_inputs(prompts=[full],**kwargs) pre=prepare_inputs(prompts=[prefix],**kwargs) ids=inp.pop('input_ids') feat=model.get_input_embeddings(ids,**inp) n=ids.shape[1]-1 if n>512: raise ValueError(f'bounded pilot length exceeded: {n}') data={'h':feat.inputs_embeds[:,:-1].astype(mx.float32), 'position_ids':feat.position_ids[...,:-1], 'labels':ids[:,1:], 'loss_mask':(mx.arange(n)[None,:]>=pre['input_ids'].shape[1]-1).astype(mx.float32), 'visual_mask':feat.visual_pos_masks[:,:-1]} if feat.deepstack_visual_embeds is not None: mask=np.array(feat.visual_pos_masks[:,:-1]) additions=[] for values in feat.deepstack_visual_embeds: dense=np.zeros(data['h'].shape,dtype=np.float32) dense[mask]=np.array(values.astype(mx.float32)) additions.append(mx.array(dense)) data['deepstack']=mx.stack(additions) mx.eval(data) mx.savez(str(Path(args.out)/'batch.npz'),**data) info={'row_id':row['id'],'image_sha256':row['image_sha256'],'tokens':n, 'supervised_tokens':float(data['loss_mask'].sum().item()),'scope':'synthetic diagram; frozen vision; one-example numerical pilot'} (Path(args.out)/'batch.json').write_text(json.dumps(info,indent=2)) emit('prepared',**info) def checkpoint(out,stage,opt,step): folder=out/f'step-{step:04}' if folder.exists(): committed=json.loads((out/'committed.json').read_text())['step'] if (out/'committed.json').exists() else 0 if step<=committed:raise FileExistsError('refusing to overwrite committed checkpoint') folder.rename(out/(folder.name+'-uncommitted-'+str(time.time_ns()))) folder.mkdir(exist_ok=False) save_tree(folder/'adapter.safetensors',stage.trainable_parameters()) save_tree(folder/'optimizer.safetensors',opt.state) (folder/'meta.json').write_text(json.dumps({'step':step,'optimizer':'Adam','lr':1e-4,'rank':4,'alpha':8})) return folder def stats(): return {'mlx_peak_bytes':mx.get_peak_memory(),'process_maxrss_bytes':resource.getrusage(resource.RUSAGE_SELF).ru_maxrss} def run(args): out=Path(args.out);out.mkdir(parents=True,exist_ok=True) mx.set_memory_limit(8*1024**3);mx.set_cache_limit(128*1024**2) if args.mode=='prepare':return prepare(args) model,processor,a,b=make_stages(args.base) batch=mx.load(args.batch) if args.mode=='reference': deep = [mx.array(np.array(v)[np.array(batch['visual_mask'])]) for v in batch['deepstack']] if 'deepstack' in batch else None native_h = model.language_model.model( mx.zeros(batch['labels'].shape,dtype=mx.int32),inputs_embeds=batch['h'], mask='causal',position_ids=batch['position_ids'], visual_pos_masks=batch['visual_mask'],deepstack_visual_embeds=deep) native_loss=loss(model.language_model.lm_head(native_h),batch) mx.eval(native_loss);native_value=native_loss.item() del native_h,native_loss,model,processor class Whole(nn.Module): def __init__(self):super().__init__();self.a=a;self.b=b def __call__(self):return loss(self.b(self.a(batch['h'],batch),batch),batch) w=Whole(); opt=optim.Adam(1e-4) for step in range(args.steps): value,grads=nn.value_and_grad(w,lambda m:m())(w) mx.eval(value,grads) if step==0: delta=abs(value.item()-native_value) emit('native-forward-parity',native_loss=native_value,staged_loss=value.item(),absolute_error=delta) if delta>1e-5:raise ValueError('native/staged forward mismatch') save_tree(out/'first-grad-a.safetensors',grads['a']) save_tree(out/'first-grad-b.safetensors',grads['b']) opt.update(w,grads);mx.eval(w.parameters(),opt.state) emit('reference-step',step=step,loss=value.item(),**stats()) save_tree(out/'final-a.safetensors',a.trainable_parameters()) save_tree(out/'final-b.safetensors',b.trainable_parameters()) (out/'result.json').write_text(json.dumps({'steps':args.steps,'last_loss':value.item(),**stats()},indent=2)) return del model,processor stage=a if args.mode=='first' else b if args.mode=='first':del b else:del a opt=optim.Adam(1e-4);start=args.resume if start: folder=out/f'step-{start:04}' stage.update(load_tree(folder/'adapter.safetensors')) opt.state=load_tree(folder/'optimizer.safetensors') mx.eval(stage.parameters(),opt.state) config_id=hashlib.sha256(Path(args.batch).read_bytes()).hexdigest() if args.mode=='last': listener=socket.socket();listener.setsockopt(socket.SOL_SOCKET,socket.SO_REUSEADDR,1) listener.bind((args.host,args.port));listener.listen(1);listener.settimeout(300) emit('listening',host=args.host,port=args.port,start=start) s,addr=listener.accept();listener.close();emit('connected',peer=addr[0]) else:s=socket.create_connection((args.host,args.port),timeout=120) s.settimeout(300) send(s,{'protocol':1,'start':start,'batch_sha256':config_id}) peer,_=recv(s) if peer!={'protocol':1,'start':start,'batch_sha256':config_id}:raise ValueError('peer checkpoint/batch mismatch') try: for step in range(start,start+args.steps): begin=time.monotonic() if args.mode=='first': h=stage(batch['h'],batch);mx.eval(h) sent=send(s,{'op':'forward','step':step},{'h':h.astype(mx.float32)}) if step==args.disconnect_step: raise ConnectionAbortedError('intentional transport interruption after forward') meta,data=recv(s) if meta.get('op')!='gradient' or meta.get('step')!=step:raise ValueError('invalid gradient reply') def objective(params): stage.update(params) return (stage(batch['h'],batch).astype(mx.float32)*data['dh']).sum() _,grads=mx.value_and_grad(objective)(stage.trainable_parameters()) value=meta['loss'] else: meta,data=recv(s) if meta!={'op':'forward','step':step}:raise ValueError('invalid forward request') def objective(params,h): stage.update(params);return loss(stage(h,batch),batch) value,(grads,dh)=mx.value_and_grad(objective,argnums=(0,1))(stage.trainable_parameters(),data['h']) mx.eval(value,grads,dh);value=float(value.item()) sent=send(s,{'op':'gradient','step':step,'loss':value},{'dh':dh.astype(mx.float32)}) mx.eval(grads) if step==0:save_tree(out/'first-grad.safetensors',grads) if not np.isfinite(value):raise ValueError('nonfinite loss') if not all(bool(mx.all(mx.isfinite(v)).item()) for v in flat(grads).values()):raise ValueError('nonfinite gradient') opt.update(stage,grads);mx.eval(stage.parameters(),opt.state) checkpoint(out,stage,opt,step+1) send(s,{'op':'saved','step':step+1});ack,_=recv(s) if ack!={'op':'saved','step':step+1}:raise ValueError('checkpoint acknowledgement mismatch') (out/'committed.json').write_text(json.dumps({'step':step+1,'batch_sha256':config_id})) emit('distributed-step',rank=args.mode,step=step,loss=value,seconds=time.monotonic()-begin,sent_bytes=sent,**stats()) mx.clear_cache() save_tree(out/'final.safetensors',stage.trainable_parameters()) (out/f'result-{start}.json').write_text(json.dumps({'start':start,'steps':args.steps,'last_loss':value,**stats()},indent=2)) finally:s.close() if __name__=='__main__': p=argparse.ArgumentParser();p.add_argument('mode',choices=['prepare','reference','first','last']) p.add_argument('--base',required=True);p.add_argument('--out',required=True) p.add_argument('--dataset');p.add_argument('--batch');p.add_argument('--steps',type=int,default=2) p.add_argument('--resume',type=int,default=0);p.add_argument('--host',default='127.0.0.1');p.add_argument('--port',type=int,default=29471) p.add_argument('--disconnect-step',type=int,default=-1) run(p.parse_args())