(version:str, pretrained:str, devices:list[str])
| 163 | return (q.matmul(k.transpose(-2,-1), dtype=dtypes.float32) / math.sqrt(q.shape[-1])).softmax(-1).cast(q.dtype) @ v |
| 164 | |
| 165 | def 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 |
searching dependent graphs…