Brainmu-SpikeCamera / src /ui /engine.py
sunbaby's picture
Upload 69 files
4719196
Raw History Blame Contribute Delete
4.34 kB
"""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)