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()