MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / load_diffusion_model

Method load_diffusion_model

models/auraflow.py:58–76  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

56 return getattr(self.diffusers_pipeline, name)
57
58 def load_diffusion_model(self):
59 dtype = self.model_config['dtype']
60 transformer_dtype = self.model_config.get('transformer_dtype', dtype)
61
62 with open('configs/auraflow/transformer_config.json') as f:
63 json_config = json.load(f)
64 with init_empty_weights():
65 transformer = diffusers.AuraFlowTransformer2DModel.from_config(json_config)
66 state_dict = {
67 k.replace('model.', ''): v
68 for k, v in load_state_dict(self.model_config['transformer_path']).items()
69 }
70 state_dict = diffusers.loaders.single_file_utils.convert_auraflow_transformer_checkpoint_to_diffusers(state_dict)
71 for key, tensor in state_dict.items():
72 dtype_to_use = dtype if any(keyword in key for keyword in KEEP_IN_HIGH_PRECISION) or tensor.ndim == 1 else transformer_dtype
73 set_module_tensor_to_device(transformer, key, device='cpu', dtype=dtype_to_use, value=tensor)
74
75 transformer.train()
76 self.diffusers_pipeline.transformer = transformer
77
78 def get_vae(self):
79 return self.vae

Callers

nothing calls this directly

Calls 2

load_state_dictFunction · 0.90
getMethod · 0.80

Tested by

no test coverage detected