NitroGen-RTX2060-ONNX / scripts /serve_onnx.py
patdev's picture
Reduce RTX 2060 VRAM peak and pin NitroGen runtime
c61e4fd verified
Raw
History Blame Contribute Delete
6.62 kB
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()