Views
No views yet
1import torch
2from diffusers import StableDiffusion3Pipeline
3
4pipe = StableDiffusion3Pipeline.from_pretrained("yujiepan/stable-diffusion-3-tiny-random", torch_dtype=torch.float16)
5pipe = pipe.to("cuda")
6
7image = pipe(
8 "A cat holding a sign that says hello world",
9 negative_prompt="",
10 num_inference_steps=2,
11 guidance_scale=7.0,
12).images[0]
13image1import importlib
2
3import torch
4import transformers
5
6import diffusers
7import rich
8
9
10def get_original_model_configs(pipeline_cls: type[diffusers.DiffusionPipeline], pipeline_id: str):
11 pipeline_config: dict[str, list[str]] = pipeline_cls.load_config(pipeline_id)
12 model_configs = {}
13
14 for subfolder, import_strings in pipeline_config.items():
15 if subfolder.startswith("_"):
16 continue
17 module = importlib.import_module(".".join(import_strings[:-1]))
18 cls = getattr(module, import_strings[-1])
19 if issubclass(cls, transformers.PreTrainedModel):
20 config_class: transformers.PretrainedConfig = cls.config_class
21 config = config_class.from_pretrained(pipeline_id, subfolder=subfolder)
22 model_configs[subfolder] = config
23 elif issubclass(cls, diffusers.ModelMixin) and issubclass(cls, diffusers.ConfigMixin):
24 config = cls.load_config(pipeline_id, subfolder=subfolder)
25 model_configs[subfolder] = config
26
27 return model_configs
28
29
30def load_pipeline(pipeline_cls: type[diffusers.DiffusionPipeline], pipeline_id: str, model_configs: dict[str, dict]):
31 pipeline_config: dict[str, list[str]] = pipeline_cls.load_config(pipeline_id)
32 components = {}
33 for subfolder, import_strings in pipeline_config.items():
34 if subfolder.startswith("_"):
35 continue
36 module = importlib.import_module(".".join(import_strings[:-1]))
37 cls = getattr(module, import_strings[-1])
38 print(f"Loading:", ".".join(import_strings))
39 if issubclass(cls, transformers.PreTrainedModel):
40 config = model_configs[subfolder]
41 component = cls(config)
42 elif issubclass(cls, transformers.PreTrainedTokenizerBase):
43 component = cls.from_pretrained(pipeline_id, subfolder=subfolder)
44 elif issubclass(cls, diffusers.ModelMixin) and issubclass(cls, diffusers.ConfigMixin):
45 config = model_configs[subfolder]
46 component = cls.from_config(config)
47 elif issubclass(cls, diffusers.SchedulerMixin) and issubclass(cls, diffusers.ConfigMixin):
48 component = cls.from_pretrained(pipeline_id, subfolder=subfolder)
49 else:
50 raise (f"unknown {subfolder}: {import_strings}")
51 components[subfolder] = component
52 pipeline = pipeline_cls(**components)
53 return pipeline
54
55
56def get_pipeline():
57 torch.manual_seed(42)
58 pipeline_id = "stabilityai/stable-diffusion-3-medium-diffusers"
59 pipeline_cls = diffusers.StableDiffusion3Pipeline
60 model_configs = get_original_model_configs(pipeline_cls, pipeline_id)
61 rich.print(model_configs)
62
63 HIDDEN_SIZE = 8
64
65 model_configs["text_encoder"].hidden_size = HIDDEN_SIZE
66 model_configs["text_encoder"].intermediate_size = HIDDEN_SIZE * 2
67 model_configs["text_encoder"].num_attention_heads = 2
68 model_configs["text_encoder"].num_hidden_layers = 2
69 model_configs["text_encoder"].projection_dim = HIDDEN_SIZE
70
71 model_configs["text_encoder_2"].hidden_size = HIDDEN_SIZE
72 model_configs["text_encoder_2"].intermediate_size = HIDDEN_SIZE * 2
73 model_configs["text_encoder_2"].num_attention_heads = 2
74 model_configs["text_encoder_2"].num_hidden_layers = 2
75 model_configs["text_encoder_2"].projection_dim = HIDDEN_SIZE
76
77 model_configs["text_encoder_3"].d_model = HIDDEN_SIZE
78 model_configs["text_encoder_3"].d_ff = HIDDEN_SIZE * 2
79 model_configs["text_encoder_3"].d_kv = HIDDEN_SIZE // 2
80 model_configs["text_encoder_3"].num_heads = 2
81 model_configs["text_encoder_3"].num_layers = 2
82
83 model_configs["transformer"]["num_layers"] = 2
84 model_configs["transformer"]["num_attention_heads"] = 2
85 model_configs["transformer"]["attention_head_dim"] = HIDDEN_SIZE // 2
86 model_configs["transformer"]["pooled_projection_dim"] = HIDDEN_SIZE * 2
87 model_configs["transformer"]["joint_attention_dim"] = HIDDEN_SIZE
88 model_configs["transformer"]["caption_projection_dim"] = HIDDEN_SIZE
89
90 model_configs["vae"]["layers_per_block"] = 1
91 model_configs["vae"]["block_out_channels"] = [HIDDEN_SIZE] * 4
92 model_configs["vae"]["norm_num_groups"] = 2
93 model_configs["vae"]["latent_channels"] = 16
94
95 pipeline = load_pipeline(pipeline_cls, pipeline_id, model_configs)
96 return pipeline
97
98
99pipeline = get_pipeline()
100image = pipeline(
101 "hello world",
102 negative_prompt="runtime error",
103 num_inference_steps=2,
104 guidance_scale=7.0,
105).images[0]
106
107
108pipeline = pipeline.to(torch.float16)
109pipeline.save_pretrained("/tmp/stable-diffusion-3-tiny-random")
110pipeline.push_to_hub("yujiepan/stable-diffusion-3-tiny-random")