"""Same frozen inference settings and metrics as the migrated reference.""" from pathlib import Path from types import SimpleNamespace import sys, os, json, importlib.util import cv2, numpy as np, torch from PIL import Image from accelerate import init_empty_weights, load_checkpoint_and_dispatch ROOT=Path(__file__).resolve().parents[1] PROJECT=ROOT sys.path.insert(0,str(PROJECT/'code')) from infer_brainmu_lora import inject_lora, metrics, set_seed from train_recon import Net, load_dat from project_config import load_config,load_frontend_weights class Engine: def load(self,model_path,adapter,frontend): self.project_config=load_config();self.settings=self.project_config['generation']['inference'] repo=ROOT/'vendor/Brainmu' sys.path.insert(0,str(repo)); os.chdir(repo) from data.data_utils import add_special_tokens from data.transforms import ImageTransform from inferencer import InterleaveInferencer from modeling.autoencoder import load_ae from modeling.brainmu import Brainmu,BrainmuConfig,Qwen2Config,Qwen2ForCausalLM,SiglipVisionConfig,SiglipVisionModel from modeling.qwen2 import Qwen2Tokenizer a=SimpleNamespace(model_path=Path(model_path),adapter=Path(adapter),adapter_config=Path(adapter).with_name('adapter_config.json')) model_path=a.model_path.resolve(); device=torch.device('cuda:0'); torch.cuda.set_device(device); torch.backends.cuda.matmul.allow_tf32=True llm=Qwen2Config.from_json_file(str(model_path/'llm_config.json')); llm.qk_norm=True; llm.tie_word_embeddings=False; llm.layer_module='Qwen2MoTDecoderLayer' vit=SiglipVisionConfig.from_json_file(str(model_path/'vit_config.json')); vit.rope=False; vit.num_hidden_layers-=1 vae,vae_cfg=load_ae(local_path=str(model_path/'ae.safetensors')) cfg=BrainmuConfig(visual_gen=True,visual_und=True,llm_config=llm,vit_config=vit,vae_config=vae_cfg,vit_max_num_patch_per_side=70,connector_act='gelu_pytorch_tanh',latent_patch_size=2,max_latent_size=64,timestep_shift=1.0) with init_empty_weights(): lm=Qwen2ForCausalLM(llm); vm=SiglipVisionModel(vit); model=Brainmu(lm,vm,cfg); model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit,meta=True) print('LOAD_BASE_BEGIN',flush=True) model=load_checkpoint_and_dispatch(model,checkpoint=str(model_path/'ema.safetensors'),device_map={'':0},dtype=torch.bfloat16,force_hooks=True) model.requires_grad_(False).eval(); vae=vae.to(device=device,dtype=torch.float32).eval().requires_grad_(False) oe,od=vae.encode,vae.decode; vae.encode=lambda x:oe(x.to(device=device,dtype=torch.float32)); vae.decode=lambda z:od(z.to(device=device,dtype=torch.float32)) spec=json.load(open(a.adapter_config)); adapters=inject_lora(model,spec,a.adapter); model.eval(); print(f'LORA_LOADED modules={len(adapters)} adapter={a.adapter}',flush=True) tok=Qwen2Tokenizer.from_pretrained(str(model_path)); tok,new_ids,_=add_special_tokens(tok) infer=InterleaveInferencer(model,vae,tok,ImageTransform(400,256,16),ImageTransform(392,252,14),new_ids) self.infer=infer self.net=Net().cuda().eval() self.net.load_state_dict(load_frontend_weights(frontend)) @torch.inference_mode() def predict(self,dat,gt,index,out,prompt): # Conditioning export precedes Brainmu; reset the same seed per sorted sample. x=torch.from_numpy(load_dat(str(dat)))[None].cuda() p=self.net(x).clamp(0,1).float().cpu().numpy()[0,0] condition=out/'condition'/(dat.stem+'.png') Image.fromarray(np.round(p*255).astype(np.uint8)).save(condition) set_seed(self.settings['seed']+index) inp=Image.open(condition).convert('RGB') pred=self.infer.interleave_inference([inp,prompt],think=False,understanding_output=False,cfg_text_scale=self.settings['cfg_text_scale'],cfg_img_scale=self.settings['cfg_img_scale'],cfg_interval=self.settings['cfg_interval'],timestep_shift=self.settings['timestep_shift'],num_timesteps=self.settings['steps'],cfg_renorm_min=self.settings['cfg_renorm_min'],cfg_renorm_type=self.settings['cfg_renorm_type'])[-1] output=out/'prediction'/(dat.stem+'.png');pred.save(output) result=metrics(pred,Image.open(gt).convert('RGB')) return dict(id=dat.stem,index=index,seed=self.settings['seed']+index,prompt=prompt,spike=str(dat),gt=str(gt),condition=str(condition),prediction=str(output),**result)