MCPcopy Create free account
hub / github.com/tinygrad/tinygrad / init_stable_diffusion

Function init_stable_diffusion

examples/mlperf/initializers.py:165–189  ·  view source on GitHub ↗
(version:str, pretrained:str, devices:list[str])

Source from the content-addressed store, hash-verified

163 return (q.matmul(k.transpose(-2,-1), dtype=dtypes.float32) / math.sqrt(q.shape[-1])).softmax(-1).cast(q.dtype) @ v
164
165def init_stable_diffusion(version:str, pretrained:str, devices:list[str]):
166 from examples.stable_diffusion import StableDiffusion
167 from tinygrad.nn.state import safe_load, safe_save, load_state_dict, get_state_dict
168 from tempfile import TemporaryDirectory
169 model = StableDiffusion(version=version, pretrained=pretrained)
170 unet:UNetModel = model.model.diffusion_model
171
172 # this prevents extra consumption of memory, enabling much larger BS
173 Tensor.realize(*get_parameters(unet))
174 with TemporaryDirectory(prefix="unet_init") as tmp:
175 safe_save(get_state_dict(unet), init_fn:=f"{tmp}/init_model.safetensors")
176 load_state_dict(unet, safe_load(init_fn))
177
178 sqrt_alphas_cumprod = model.alphas_cumprod.sqrt().realize()
179 sqrt_one_minus_alphas_cumprod = (1 - model.alphas_cumprod).sqrt().realize()
180
181 if len(devices) > 1:
182 to_move = [sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod]
183 if version == "v2-mlperf-train": to_move += get_parameters(unet) + get_parameters(model.cond_stage_model)
184 for p in to_move:
185 p.to_(devices)
186 with Context(BEAM=0):
187 Tensor.realize(*to_move)
188
189 return model, unet, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod

Callers 3

helper_test_initMethod · 0.90
train_stable_diffusionFunction · 0.90
eval_stable_diffusionFunction · 0.90

Calls 10

StableDiffusionClass · 0.90
get_parametersFunction · 0.90
safe_saveFunction · 0.90
get_state_dictFunction · 0.90
load_state_dictFunction · 0.90
safe_loadFunction · 0.90
ContextClass · 0.90
realizeMethod · 0.80
sqrtMethod · 0.80
to_Method · 0.80

Tested by 1

helper_test_initMethod · 0.72

Used in the wild real call sites across dependent graphs

searching dependent graphs…