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 argparse, gc, pickle, sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import zmq | |
| from huggingface_hub import hf_hub_download | |
| from transformers import AutoImageProcessor | |
| ROOT = Path(__file__).resolve().parents[1] | |
| VENDOR = ROOT / '.vendor' / 'NitroGen' | |
| if VENDOR.exists(): | |
| sys.path.insert(0, str(VENDOR)) | |
| sys.path.insert(0, str(ROOT)) | |
| from src.ort_dit import OrtDitModule | |
| 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 | |
| def load_hybrid_model(checkpoint_path: str): | |
| """Load all upstream NitroGen weights except the DiT, which is provided by ONNX. | |
| This avoids moving the ~122M-parameter PyTorch DiT onto a 6 GB GPU only to | |
| delete it immediately afterwards. | |
| """ | |
| checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=False) | |
| ckpt_config = CkptConfig.model_validate(checkpoint['ckpt_config']) | |
| model_cfg = ckpt_config.model_cfg | |
| tokenizer_cfg = ckpt_config.tokenizer_cfg | |
| if not isinstance(model_cfg, NitroGen_Config): | |
| raise RuntimeError(f'Unsupported NitroGen config: {type(model_cfg)}') | |
| if not isinstance(tokenizer_cfg, NitrogenTokenizerConfig): | |
| raise RuntimeError(f'Unsupported tokenizer config: {type(tokenizer_cfg)}') | |
| img_proc = AutoImageProcessor.from_pretrained(model_cfg.vision_encoder_name) | |
| tokenizer_cfg.training = False | |
| tokenizer = NitrogenTokenizer(tokenizer_cfg) | |
| game_mapping = tokenizer.game_mapping | |
| model = NitroGen(config=model_cfg, game_mapping=game_mapping) | |
| # Remove randomly initialized DiT before loading/transferring weights. | |
| model.model = torch.nn.Identity() | |
| non_dit_state = {k: v for k, v in checkpoint['model'].items() if not k.startswith('model.')} | |
| missing, unexpected = model.load_state_dict(non_dit_state, strict=False) | |
| if unexpected: | |
| raise RuntimeError(f'Unexpected checkpoint keys: {unexpected[:8]}') | |
| # Missing keys are only acceptable if upstream added buffers to the removed DiT. | |
| bad_missing = [k for k in missing if not k.startswith('model.')] | |
| if bad_missing: | |
| raise RuntimeError(f'Missing non-DiT checkpoint keys: {bad_missing[:8]}') | |
| del checkpoint, non_dit_state | |
| gc.collect() | |
| model.eval().half().to('cuda') | |
| return model, tokenizer, img_proc, ckpt_config, game_mapping, 1 | |
| class TuringFp16Session(InferenceSession): | |
| """NitroGen session tuned for Turing GPUs (RTX 20xx): FP16 instead of BF16.""" | |
| def _predict_flowmatching(self, pixel_values, action_tensors): | |
| available_frames = len(self.obs_buffer) | |
| frames = torch.zeros( | |
| (self.max_buffer_size, *pixel_values.shape[1:]), | |
| dtype=pixel_values.dtype, | |
| 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): | |
| data[k] = v.unsqueeze(0).to('cuda') | |
| 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): | |
| if self.cfg_scale == 1.0: | |
| out = self.model.get_action(tok_hist, old_layout=self.old_layout) | |
| else: | |
| out = self.model.get_action_with_cfg(tok_hist, tok_nohist, cfg_scale=self.cfg_scale) | |
| return self.tokenizer.decode(out) | |
| def main(): | |
| ap = argparse.ArgumentParser(description='NitroGen ONNX FP16 server for RTX 2060/Turing') | |
| ap.add_argument('--ckpt', default=None) | |
| ap.add_argument('--onnx', default=None) | |
| ap.add_argument('--port', type=int, default=5555) | |
| ap.add_argument('--steps', type=int, choices=[4, 8, 16], default=4) | |
| ap.add_argument('--cfg', type=float, default=1.0) | |
| ap.add_argument('--ctx', type=int, default=1) | |
| ap.add_argument('--no-trt', action='store_true') | |
| args = ap.parse_args() | |
| if not torch.cuda.is_available(): | |
| raise SystemExit('CUDA GPU required for gameplay server') | |
| ckpt = args.ckpt or hf_hub_download('nvidia/NitroGen', 'ng.pt') | |
| onnx = args.onnx or hf_hub_download('patdev/NitroGen-RTX2060-ONNX', 'onnx/dit_fp16.onnx') | |
| model, tokenizer, img_proc, ckpt_cfg, game_mapping, ratio = load_hybrid_model(ckpt) | |
| model.num_inference_timesteps = args.steps | |
| model.model = OrtDitModule( | |
| onnx, | |
| prefer_tensorrt=not args.no_trt, | |
| cache_dir=ROOT / '.ort-cache', | |
| ) | |
| gc.collect() | |
| torch.cuda.empty_cache() | |
| session = TuringFp16Session( | |
| model, ckpt, tokenizer, img_proc, ckpt_cfg, game_mapping, | |
| None, False, args.cfg, ratio, args.ctx, | |
| ) | |
| print( | |
| f'NitroGen ONNX ready | GPU={torch.cuda.get_device_name(0)} ' | |
| f'| EP={model.model.provider} | steps={args.steps}' | |
| ) | |
| context = zmq.Context() | |
| sock = context.socket(zmq.REP) | |
| sock.bind(f'tcp://*:{args.port}') | |
| try: | |
| while True: | |
| req = pickle.loads(sock.recv()) | |
| if req['type'] == 'reset': | |
| session.reset(); resp = {'status': 'ok'} | |
| elif req['type'] == 'info': | |
| resp = {'status': 'ok', 'info': session.info() | { | |
| 'provider': model.model.provider, | |
| 'steps': args.steps, | |
| }} | |
| elif req['type'] == 'predict': | |
| resp = {'status': 'ok', 'pred': session.predict(req['image'])} | |
| else: | |
| resp = {'status': 'error', 'message': 'unknown request'} | |
| sock.send(pickle.dumps(resp)) | |
| except KeyboardInterrupt: | |
| pass | |
| finally: | |
| sock.close(); context.term() | |
| if __name__ == '__main__': | |
| main() | |