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

Method load_diffusion_model

models/chroma.py:148–168  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

146 return getattr(self.diffusers_pipeline, name)
147
148 def load_diffusion_model(self):
149 dtype = self.model_config['dtype']
150 transformer_dtype = self.model_config.get('transformer_dtype', dtype)
151 with init_empty_weights():
152 transformer = Chroma(chroma_params)
153 transformer.load_state_dict(load_state_dict(self.model_config['transformer_path']), assign=True)
154 self.diffusers_pipeline.transformer = transformer
155
156 if 'adapter' in self.config:
157 if fuse_adapters := self.config['adapter'].get('fuse_adapters', None):
158 print(f'Fusing adapters: {fuse_adapters}')
159 for fuse_adapter in fuse_adapters:
160 self.load_and_fuse_adapter(fuse_adapter['path'])
161
162 for name, p in transformer.named_parameters():
163 if not any(x in name for x in KEEP_IN_HIGH_PRECISION):
164 p.data = p.data.to(transformer_dtype)
165
166 self.transformer.train()
167 for name, p in self.transformer.named_parameters():
168 p.original_name = name
169
170 def get_vae(self):
171 return self.vae

Callers

nothing calls this directly

Calls 5

load_state_dictFunction · 0.90
getMethod · 0.80
load_state_dictMethod · 0.45
load_and_fuse_adapterMethod · 0.45
toMethod · 0.45

Tested by

no test coverage detected