CoolFace
Modelpublic

BAAI/Brainmu-Spike

sourceHugging Faceapache-2.0updated 10d agoView on Hugging Face
11likes10kdownloads
engine.py59 linesDownload Raw Back to ui
1"""Same frozen inference settings and metrics as the migrated reference."""2from pathlib import Path3from types import SimpleNamespace4import sys, os, json, importlib.util5import cv2, numpy as np, torch6from PIL import Image7from accelerate import init_empty_weights, load_checkpoint_and_dispatch8ROOT=Path(__file__).resolve().parents[1]9PROJECT=ROOT10sys.path.insert(0,str(PROJECT/'code'))11from infer_brainmu_lora import inject_lora, metrics, set_seed12from train_recon import Net, load_dat13from project_config import load_config,load_frontend_weights14 15class Engine:16 def load(self,model_path,adapter,frontend):17    self.project_config=load_config();self.settings=self.project_config['generation']['inference']18    repo=ROOT/'vendor/Brainmu'19    sys.path.insert(0,str(repo)); os.chdir(repo)20    from data.data_utils import add_special_tokens21    from data.transforms import ImageTransform22    from inferencer import InterleaveInferencer23    from modeling.autoencoder import load_ae24    from modeling.brainmu import Brainmu,BrainmuConfig,Qwen2Config,Qwen2ForCausalLM,SiglipVisionConfig,SiglipVisionModel25    from modeling.qwen2 import Qwen2Tokenizer26    a=SimpleNamespace(model_path=Path(model_path),adapter=Path(adapter),adapter_config=Path(adapter).with_name('adapter_config.json'))27    model_path=a.model_path.resolve(); device=torch.device('cuda:0'); torch.cuda.set_device(device); torch.backends.cuda.matmul.allow_tf32=True28    llm=Qwen2Config.from_json_file(str(model_path/'llm_config.json')); llm.qk_norm=True; llm.tie_word_embeddings=False; llm.layer_module='Qwen2MoTDecoderLayer'29    vit=SiglipVisionConfig.from_json_file(str(model_path/'vit_config.json')); vit.rope=False; vit.num_hidden_layers-=130    vae,vae_cfg=load_ae(local_path=str(model_path/'ae.safetensors'))31    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)32    with init_empty_weights():33        lm=Qwen2ForCausalLM(llm); vm=SiglipVisionModel(vit); model=Brainmu(lm,vm,cfg); model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit,meta=True)34    print('LOAD_BASE_BEGIN',flush=True)35    model=load_checkpoint_and_dispatch(model,checkpoint=str(model_path/'ema.safetensors'),device_map={'':0},dtype=torch.bfloat16,force_hooks=True)36    model.requires_grad_(False).eval(); vae=vae.to(device=device,dtype=torch.float32).eval().requires_grad_(False)37    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))38    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)39    tok=Qwen2Tokenizer.from_pretrained(str(model_path)); tok,new_ids,_=add_special_tokens(tok)40    infer=InterleaveInferencer(model,vae,tok,ImageTransform(400,256,16),ImageTransform(392,252,14),new_ids)41 42    self.infer=infer43    self.net=Net().cuda().eval()44    self.net.load_state_dict(load_frontend_weights(frontend))45 46 @torch.inference_mode()47 def predict(self,dat,gt,index,out,prompt):48    # Conditioning export precedes Brainmu; reset the same seed per sorted sample.49    x=torch.from_numpy(load_dat(str(dat)))[None].cuda()50    p=self.net(x).clamp(0,1).float().cpu().numpy()[0,0]51    condition=out/'condition'/(dat.stem+'.png')52    Image.fromarray(np.round(p*255).astype(np.uint8)).save(condition)53    set_seed(self.settings['seed']+index)54    inp=Image.open(condition).convert('RGB')55    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]56    output=out/'prediction'/(dat.stem+'.png');pred.save(output)57    result=metrics(pred,Image.open(gt).convert('RGB'))58    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)59