File size: 8,017 Bytes
a806943
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7a12b11
 
a806943
 
 
 
 
 
 
 
7a12b11
 
 
a806943
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1226aef
 
 
 
 
 
 
 
 
a806943
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
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()