Instructions to use BAAI/Brainmu-SpikeCamera with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BAAI/Brainmu-SpikeCamera with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 4,335 Bytes
4719196 | 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 | """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)
|