Instructions to use patdev/NitroGen-RTX2060-ONNX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- TensorRT
How to use patdev/NitroGen-RTX2060-ONNX with TensorRT:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
| 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() | |