from __future__ import annotations import asyncio, gc, os, sys, threading, time from difflib import SequenceMatcher from pathlib import Path from typing import Any import numpy as np import torch from PIL import Image from transformers import AutoImageProcessor VENDOR = Path(os.environ.get('NITROGEN_VENDOR','/opt/NitroGen')) if str(VENDOR) not in sys.path: sys.path.insert(0,str(VENDOR)) from nitrogen.cfg import CkptConfig from nitrogen.flow_matching_transformer.nitrogen import NitroGen, NitroGen_Config from nitrogen.mm_tokenizers import NitrogenTokenizerConfig, NitrogenTokenizer from nitrogen.inference_session import InferenceSession from ort_dit import OrtDitModule MODEL_REPO = Path(os.environ.get('NITROGEN_ONNX_ROOT','/models/nitrogen-onnx')) UPSTREAM_REPO = Path(os.environ.get('NITROGEN_UPSTREAM_ROOT','/models/nitrogen-upstream')) DATA_ROOT = Path(os.environ.get('DATA_ROOT','/data')) CACHE = DATA_ROOT/'ort-cache'; CACHE.mkdir(parents=True,exist_ok=True) class TuringFp16Session(InferenceSession): def _predict_flowmatching(self,pixel_values,action_tensors): available_frames=len(self.obs_buffer) pixel_values=pixel_values.to(device='cuda',dtype=torch.float16) frames=torch.zeros((self.max_buffer_size,*pixel_values.shape[1:]),dtype=torch.float16,device='cuda') frames[-available_frames:]=pixel_values dropped_frames=torch.zeros((self.max_buffer_size,),dtype=torch.bool,device='cuda') dropped_frames[:self.max_buffer_size-available_frames]=True tok_hist=self.tokenizer.encode({'frames':frames,'dropped_frames':dropped_frames,'game':self.selected_game}) frame_mask=torch.ones((self.max_buffer_size,),dtype=torch.bool,device='cuda'); frame_mask[-1]=False tok_nohist=self.tokenizer.encode({'frames':frames,'dropped_frames':frame_mask,'game':None}) for data in (tok_hist,tok_nohist): for k,v in list(data.items()): if isinstance(v,torch.Tensor): t=v.unsqueeze(0).to('cuda') data[k]=t.to(torch.float16) if t.is_floating_point() else t elif isinstance(v,np.ndarray): data[k]=torch.tensor(v,device='cuda').unsqueeze(0) else: data[k]=[v] with torch.inference_mode(),torch.autocast(device_type='cuda',dtype=torch.float16): out=self.model.get_action(tok_hist,old_layout=self.old_layout) if self.cfg_scale==1.0 else self.model.get_action_with_cfg(tok_hist,tok_nohist,cfg_scale=self.cfg_scale) return self.tokenizer.decode(out) class NitroGenRuntime: def __init__(self): self.state='unloaded'; self.error=None; self.loaded_at=None; self.provider=None self.model=None; self.tokenizer=None; self.img_proc=None; self.ckpt_cfg=None; self.game_mapping=None; self.ckpt_path=None self.steps=int(os.environ.get('NITROGEN_STEPS','4')); self.cfg=float(os.environ.get('NITROGEN_CFG','1.0')) self.infer_lock=threading.Lock(); self.load_lock=asyncio.Lock() def _find_file(self,root:Path,name:str)->Path: candidates=[root/name,*root.glob(f'**/{name}')] for p in candidates: if p.exists(): return p raise FileNotFoundError(f'{name} not found under {root}') def _load_sync(self): if not torch.cuda.is_available(): raise RuntimeError('CUDA GPU required') self.state='loading'; self.error=None ckpt=self._find_file(UPSTREAM_REPO,'ng.pt'); onnx=self._find_file(MODEL_REPO,'dit_fp16.onnx') checkpoint=torch.load(str(ckpt),map_location='cpu',weights_only=False) cfg=CkptConfig.model_validate(checkpoint['ckpt_config']); model_cfg=cfg.model_cfg; tok_cfg=cfg.tokenizer_cfg if not isinstance(model_cfg,NitroGen_Config) or not isinstance(tok_cfg,NitrogenTokenizerConfig): raise RuntimeError('Unsupported NitroGen checkpoint config') img_proc=AutoImageProcessor.from_pretrained(model_cfg.vision_encoder_name) tok_cfg.training=False; tokenizer=NitrogenTokenizer(tok_cfg); game_mapping=tokenizer.game_mapping model=NitroGen(config=model_cfg,game_mapping=game_mapping); model.model=torch.nn.Identity() non_dit={k:v for k,v in checkpoint['model'].items() if not k.startswith('model.')} missing,unexpected=model.load_state_dict(non_dit,strict=False) bad=[k for k in missing if not k.startswith('model.')] if unexpected or bad: raise RuntimeError(f'Checkpoint mismatch unexpected={unexpected[:4]} missing={bad[:4]}') del checkpoint,non_dit; gc.collect() model.eval().half().to('cuda'); model.num_inference_timesteps=self.steps # NitroGen upstream uses vision.dtype to allocate sa_embs, while the Turing # action encoder runs FP16. SigLIP can return FP32 under autocast, causing # masked_scatter_(Float <- Half). Normalize both branches at the source. _prepare_input_embs = model.prepare_input_embs def _prepare_input_embs_fp16(vl_token_ids, sa_token_ids, vision, action, dropped_images, game_ids=None): vision = vision.to(dtype=torch.float16) action = action.to(dtype=torch.float16) return _prepare_input_embs(vl_token_ids, sa_token_ids, vision, action, dropped_images, game_ids=game_ids) model.prepare_input_embs = _prepare_input_embs_fp16 model.model=OrtDitModule(onnx,prefer_tensorrt=True,cache_dir=CACHE) gc.collect(); torch.cuda.empty_cache() self.model,self.tokenizer,self.img_proc,self.ckpt_cfg,self.game_mapping=model,tokenizer,img_proc,cfg,game_mapping self.ckpt_path=str(ckpt); self.provider=model.model.provider; self.loaded_at=time.time(); self.state='ready' async def ensure_loaded(self): if self.state=='ready': return async with self.load_lock: if self.state=='ready': return try: await asyncio.to_thread(self._load_sync) except Exception as e: self.state='error'; self.error=str(e); raise def match_game(self,title:str|None)->tuple[str|None,float]: if not title or not self.game_mapping: return None,0.0 q=''.join(ch.lower() if ch.isalnum() else ' ' for ch in title); q=' '.join(q.split()) best=None; score=0.0 for name in self.game_mapping: n=''.join(ch.lower() if ch.isalnum() else ' ' for ch in name); n=' '.join(n.split()) s=SequenceMatcher(None,q,n).ratio() if q in n or n in q: s=max(s,0.82) if s>score: best,score=name,s return (best,score) if score>=0.50 else (None,score) async def create_session(self,title:str|None=None,game_override:str|None=None,context:int=1): await self.ensure_loaded() selected=game_override or self.match_game(title)[0] return TuringFp16Session(self.model,self.ckpt_path,self.tokenizer,self.img_proc,self.ckpt_cfg,self.game_mapping,selected,False,self.cfg,1,context) async def predict(self,session:TuringFp16Session,image:Image.Image)->dict[str,Any]: await self.ensure_loaded() def run(): with self.infer_lock: t=time.perf_counter(); pred=session.predict(image); dt=time.perf_counter()-t return {'j_left':np.asarray(pred['j_left']).tolist(),'j_right':np.asarray(pred['j_right']).tolist(),'buttons':np.asarray(pred['buttons']).tolist(),'latency_ms':round(dt*1000,1)} return await asyncio.to_thread(run) def info(self): gpu=torch.cuda.get_device_name(0) if torch.cuda.is_available() else None mem=None if torch.cuda.is_available(): mem={'allocated_mb':round(torch.cuda.memory_allocated()/1048576,1),'reserved_mb':round(torch.cuda.memory_reserved()/1048576,1),'total_mb':round(torch.cuda.get_device_properties(0).total_memory/1048576,1)} return {'state':self.state,'error':self.error,'provider':self.provider,'steps':self.steps,'cfg':self.cfg,'gpu':gpu,'memory':mem,'loaded_at':self.loaded_at,'game_mapping_count':len(self.game_mapping or {})} runtime=NitroGenRuntime()