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:
# Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download src/ui/engine.py from BAAI/Brainmu-SpikeCamera: direct link, hf CLI and curl.
- Browser
- Download file 4.34 kB
-
https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/ui/engine.py
- Command line
-
hf download hf://BAAI/Brainmu-SpikeCamera/src/ui/engine.py
-
curl -L -o engine.py https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/ui/engine.py
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)) | |
| 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) | |