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)