justinchuby's picture
Add files using upload-large-folder tool
9afed76 verified
Raw
History Blame Contribute Delete
4.8 kB
from __future__ import annotations
import json, os, platform, subprocess, threading, time
from pathlib import Path
import numpy as np
import onnxruntime as ort
import psutil
from PIL import Image
from diffusers import DDIMScheduler, PNDMScheduler
from huggingface_hub import hf_hub_download
from transformers import CLIPTokenizer
ROOT=Path('/datadisks/disk1/justinchu/inference-metadata-catalogue/stable-diffusion-bk-sdm-small')
MODEL='nota-ai/bk-sdm-small'; REV='572238db7ed3a10858900803f3fc8cca53e893e0'
PROMPT='a small red fox sleeping under northern lights, detailed digital painting'
NEG=''; SEED=42; HEIGHT=512; WIDTH=512; STEPS=10; GUIDANCE=7.5
for sub, files in {'tokenizer':['merges.txt','vocab.json','tokenizer_config.json','special_tokens_map.json'], 'scheduler':['scheduler_config.json'], 'feature_extractor':['preprocessor_config.json']}.items():
dst=ROOT/sub; dst.mkdir(parents=True,exist_ok=True)
for name in files:
src=hf_hub_download(MODEL,f'{sub}/{name}',revision=REV)
(dst/name).write_bytes(Path(src).read_bytes())
def session(component):
return ort.InferenceSession(str(ROOT/component/'model.onnx'),providers=['CUDAExecutionProvider','CPUExecutionProvider'])
proc=psutil.Process(); stop=False; peak_rss=proc.memory_info().rss; peak_gpu=0
def monitor():
global peak_rss,peak_gpu
while not stop:
peak_rss=max(peak_rss,proc.memory_info().rss)
try:
out=subprocess.check_output(['nvidia-smi','--query-gpu=memory.used','--format=csv,noheader,nounits'],text=True)
peak_gpu=max(peak_gpu,int(out.splitlines()[0].strip()))
except Exception: pass
time.sleep(.1)
th=threading.Thread(target=monitor,daemon=True); th.start()
total_start=time.perf_counter(); timings={}
t=time.perf_counter(); tok=CLIPTokenizer.from_pretrained(str(ROOT/'tokenizer')); scheduler=DDIMScheduler.from_config(PNDMScheduler.from_pretrained(str(ROOT/'scheduler')).config); timings['asset_load_s']=time.perf_counter()-t
t=time.perf_counter(); text_s=session('text_encoder'); unet_s=session('unet'); vae_s=session('vae_decoder'); timings['session_load_s']=time.perf_counter()-t
def encode(text):
ids=tok(text,padding='max_length',max_length=tok.model_max_length,truncation=True,return_tensors='np').input_ids.astype(np.int64)
return text_s.run(None,{'input_ids':ids})[0].astype(np.float32)
t=time.perf_counter(); cond=encode(PROMPT); uncond=encode(NEG); embeds=np.concatenate([uncond,cond]); timings['text_encode_s']=time.perf_counter()-t
rng=np.random.default_rng(SEED); latents=rng.standard_normal((1,4,HEIGHT//8,WIDTH//8),dtype=np.float32)
scheduler.set_timesteps(STEPS); latents*=float(scheduler.init_noise_sigma)
denoise=[]
for timestep in scheduler.timesteps:
t=time.perf_counter(); model_in=np.concatenate([latents,latents]); model_in=scheduler.scale_model_input(__import__('torch').from_numpy(model_in),timestep).numpy().astype(np.float32)
t_arr=np.full((2,),float(timestep),dtype=np.float32)
pred=unet_s.run(None,{'sample':model_in,'timestep':t_arr,'encoder_hidden_states':embeds})[0]
p0,p1=np.split(pred,2); guided=p0+GUIDANCE*(p1-p0)
latents=scheduler.step(__import__('torch').from_numpy(guided),timestep,__import__('torch').from_numpy(latents)).prev_sample.numpy().astype(np.float32)
denoise.append(time.perf_counter()-t)
timings['denoise_total_s']=sum(denoise); timings['denoise_step_s']=denoise
t=time.perf_counter(); image=vae_s.run(None,{'latent_sample':(latents/0.18215).astype(np.float32)})[0]; timings['vae_decode_s']=time.perf_counter()-t
image=np.clip(image/2+0.5,0,1); image=(image[0].transpose(1,2,0)*255).round().astype(np.uint8); Image.fromarray(image).save(ROOT/'generated.png')
timings['total_s']=time.perf_counter()-total_start; stop=True; th.join(timeout=1)
meta={'model_id':MODEL,'revision':REV,'license':'creativeml-openrail-m','task':'text-to-image','prompt':PROMPT,'negative_prompt':NEG,'seed':SEED,'width':WIDTH,'height':HEIGHT,'num_inference_steps':STEPS,'guidance_scale':GUIDANCE,'scheduler':'DDIMScheduler','source_scheduler':'PNDMScheduler','runtime':{'python':platform.python_version(),'onnxruntime':ort.__version__,'diffusers':__import__('diffusers').__version__,'transformers':__import__('transformers').__version__,'numpy':np.__version__,'providers':unet_s.get_providers(),'gpu':subprocess.check_output(['nvidia-smi','--query-gpu=name,driver_version','--format=csv,noheader'],text=True).splitlines()[0]},'timings':timings,'peak_memory':{'process_rss_bytes':peak_rss,'gpu_0_used_mib':peak_gpu},'evidence':{'image':'generated.png','onnx_components':['text_encoder/model.onnx','unet/model.onnx','vae_encoder/model.onnx','vae_decoder/model.onnx']}}
(ROOT/'execution_evidence.json').write_text(json.dumps(meta,indent=2)+'\n')
print(json.dumps(meta,indent=2))