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

Method load_diffusion_model

models/sd3.py:43–55  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

41 return getattr(self.diffusers_pipeline, name)
42
43 def load_diffusion_model(self):
44 dtype = self.model_config['dtype']
45 transformer_dtype = self.model_config.get('transformer_dtype', dtype)
46 if diffusers_path := self.model_config.get('diffusers_path', None):
47 transformer = diffusers.SD3Transformer2DModel.from_pretrained(diffusers_path, torch_dtype=dtype, subfolder='transformer')
48 else:
49 raise NotImplementedError()
50
51 for name, p in transformer.named_parameters():
52 if not (any(x in name for x in KEEP_IN_HIGH_PRECISION) or p.ndim == 1):
53 p.data = p.data.to(transformer_dtype)
54
55 self.diffusers_pipeline.transformer = transformer
56
57 def get_vae(self):
58 return self.vae

Callers 3

train.pyFile · 0.45
test_ideogram4.pyFile · 0.45

Calls 3

getMethod · 0.80
from_pretrainedMethod · 0.80
toMethod · 0.45

Tested by

no test coverage detected