Skip to content

Commit ffc7baf

Browse files
committed
Address SeFi review feedback
1 parent 614828a commit ffc7baf

6 files changed

Lines changed: 269 additions & 118 deletions

File tree

docs/source/en/api/pipelines/sefi.md

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -17,14 +17,27 @@ It denoises a semantic latent stream slightly ahead of a texture latent stream,
1717
to images.
1818

1919
The public checkpoints on the Hub are gated and distributed under a non-commercial license. Accept the checkpoint
20-
license on the Hub before downloading them. Original SeFi-Image artifacts can be converted with
21-
`scripts/convert_sefi_to_diffusers.py` before loading them with [`SeFiPipeline`].
20+
license on the Hub before downloading them. Standalone Diffusers checkpoints are available for the following models:
21+
22+
| Model | Type |
23+
|---|---|
24+
| [`SeFi-Image/SeFi-Image-1B-Base-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-1B-Base-diffusers) | Base |
25+
| [`SeFi-Image/SeFi-Image-1B-turbo-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-1B-turbo-diffusers) | Turbo |
26+
| [`SeFi-Image/SeFi-Image-2B-Base-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-2B-Base-diffusers) | Base |
27+
| [`SeFi-Image/SeFi-Image-2B-turbo-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-2B-turbo-diffusers) | Turbo |
28+
| [`SeFi-Image/SeFi-Image-5B-Base-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-5B-Base-diffusers) | Base |
29+
| [`SeFi-Image/SeFi-Image-5B-turbo-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-5B-turbo-diffusers) | Turbo |
30+
| [`SeFi-Image/SeFi-Image-5B-RL-diffusers`](https://huggingface.co/SeFi-Image/SeFi-Image-5B-RL-diffusers) | RL |
31+
32+
Original SeFi-Image artifacts can also be converted with `scripts/convert_sefi_to_diffusers.py`.
2233

2334
```python
2435
import torch
2536
from diffusers import SeFiPipeline
2637

27-
pipe = SeFiPipeline.from_pretrained("./sefi-1b-base-diffusers", torch_dtype=torch.bfloat16)
38+
pipe = SeFiPipeline.from_pretrained(
39+
"SeFi-Image/SeFi-Image-1B-Base-diffusers", dtype=torch.bfloat16
40+
)
2841
pipe.to("cuda")
2942

3043
image = pipe("A red apple on a wooden table.").images[0]

scripts/convert_sefi_to_diffusers.py

Lines changed: 128 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323
from huggingface_hub import snapshot_download
2424
from safetensors.torch import load_file
2525

26-
from diffusers import __version__
26+
from diffusers import FlowMatchEulerDiscreteScheduler, __version__
2727
from diffusers.models import SeFiTransformer2DModel
2828

2929

@@ -106,6 +106,16 @@ def parse_args():
106106
parser.add_argument("--output", required=True, help="Output Diffusers checkpoint folder.")
107107
parser.add_argument("--cache-dir", default=None, help="Optional Hugging Face cache directory.")
108108
parser.add_argument("--token", default=None, help="Optional Hugging Face token for gated checkpoints.")
109+
parser.add_argument(
110+
"--source-repo-id",
111+
default=None,
112+
help="Original Hub repo id to record in the converted model card when converting a local checkpoint.",
113+
)
114+
parser.add_argument(
115+
"--target-repo-id",
116+
default=None,
117+
help="Converted Hub repo id to use in the generated model card example.",
118+
)
109119
parser.add_argument(
110120
"--variant",
111121
choices=["base", "rl", "turbo"],
@@ -144,6 +154,66 @@ def save_json(path: Path, payload: dict):
144154
handle.write("\n")
145155

146156

157+
def save_model_card(output: Path, source_repo_id: str | None, target_repo_id: str, variant: str):
158+
metadata = [
159+
"---",
160+
"license: cc-by-nc-4.0",
161+
"library_name: diffusers",
162+
"pipeline_tag: text-to-image",
163+
"gated: true",
164+
"tags:",
165+
"- sefi-image",
166+
"- semantic-first-diffusion",
167+
"- safetensors",
168+
]
169+
if source_repo_id is not None:
170+
metadata.append(f"base_model: {source_repo_id}")
171+
metadata.append("---")
172+
173+
source_link = (
174+
f"[`{source_repo_id}`](https://huggingface.co/{source_repo_id})"
175+
if source_repo_id is not None
176+
else "the original SeFi-Image checkpoint"
177+
)
178+
inference_call = (
179+
"""image = pipe(
180+
\"A red apple on a wooden table.\",
181+
num_inference_steps=4,
182+
guidance_scale=1.0,
183+
).images[0]"""
184+
if variant == "turbo"
185+
else 'image = pipe("A red apple on a wooden table.").images[0]'
186+
)
187+
card = (
188+
"\n".join(metadata)
189+
+ f"""
190+
191+
# SeFi-Image Diffusers checkpoint
192+
193+
This repository is a Diffusers-format conversion of {source_link}. The original checkpoint is not modified by the
194+
conversion. Refer to the source model card for model details, limitations, and responsible-use guidance.
195+
196+
```python
197+
import torch
198+
from diffusers import SeFiPipeline
199+
200+
pipe = SeFiPipeline.from_pretrained(
201+
\"{target_repo_id}\", dtype=torch.bfloat16
202+
).to(\"cuda\")
203+
{inference_call}
204+
image.save(\"sefi.png\")
205+
```
206+
207+
## License
208+
209+
The checkpoint is distributed under the Creative Commons Attribution-NonCommercial 4.0 International license
210+
(CC BY-NC 4.0). It is for non-commercial use only.
211+
"""
212+
)
213+
with open(output / "README.md", "w", encoding="utf-8") as handle:
214+
handle.write(card)
215+
216+
147217
def infer_variant(checkpoint: str, config: dict, explicit_variant: str | None) -> str:
148218
if explicit_variant is not None:
149219
return explicit_variant
@@ -165,9 +235,7 @@ def default_guidance_scale(variant: str) -> float:
165235

166236

167237
def texture_vae_config_path(root: Path, texture_vae_name: str) -> Path:
168-
if texture_vae_name == "sd1.5":
169-
return root / "vae" / "config.json"
170-
if texture_vae_name in {"flux1", "flux2"}:
238+
if texture_vae_name in {"sd1.5", "flux1", "flux2"}:
171239
return root / "vae" / "config.json"
172240
raise ValueError(f"Unsupported texture VAE: {texture_vae_name}")
173241

@@ -198,7 +266,6 @@ def build_transformer_config(root: Path, sefi_config: dict) -> dict:
198266

199267
transformer_config["in_channels"] = total_channels
200268
transformer_config["out_channels"] = total_channels
201-
transformer_config["text_input_dim"] = text_dim
202269
if int(transformer_config["joint_attention_dim"]) != text_dim:
203270
raise ValueError(
204271
"Text dimension mismatch: "
@@ -225,6 +292,38 @@ def load_transformer_state_dict(transformer_dir: Path) -> dict[str, torch.Tensor
225292
raise FileNotFoundError(f"No supported transformer weights found under {transformer_dir}.")
226293

227294

295+
def convert_transformer_state_dict(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
296+
converted = {}
297+
for key, value in state_dict.items():
298+
if key.startswith("backbone."):
299+
converted_key = key.removeprefix("backbone.")
300+
elif key.startswith("dual_time_embed."):
301+
converted_key = key
302+
else:
303+
raise ValueError(f"Unexpected transformer key in the original SeFi checkpoint: {key}")
304+
305+
if converted_key in converted:
306+
raise ValueError(f"Transformer key collision after conversion: {converted_key}")
307+
converted[converted_key] = value
308+
return converted
309+
310+
311+
def convert_scheduler(root: Path, output: Path):
312+
scheduler_config = load_json(root / "scheduler" / "scheduler_config.json")
313+
source_scheduler = FlowMatchEulerDiscreteScheduler.from_config(scheduler_config)
314+
scheduler_config["shift"] = 1.0
315+
scheduler_config["use_dynamic_shifting"] = False
316+
scheduler = FlowMatchEulerDiscreteScheduler.from_config(scheduler_config)
317+
318+
torch.testing.assert_close(scheduler.sigmas, source_scheduler.sigmas, rtol=0.0, atol=0.0)
319+
torch.testing.assert_close(scheduler.timesteps, source_scheduler.timesteps, rtol=0.0, atol=0.0)
320+
num_train_timesteps = int(scheduler.config.num_train_timesteps)
321+
expected_sigmas = torch.linspace(1.0, 1.0 / num_train_timesteps, num_train_timesteps)
322+
torch.testing.assert_close(scheduler.sigmas, expected_sigmas, rtol=0.0, atol=1e-7)
323+
torch.testing.assert_close(scheduler.timesteps, expected_sigmas * num_train_timesteps, rtol=0.0, atol=1e-4)
324+
scheduler.save_pretrained(output / "scheduler")
325+
326+
228327
def copy_tokenizer_files(src: Path, dst: Path):
229328
weight_patterns = {
230329
"model*.safetensors",
@@ -253,14 +352,28 @@ def main():
253352
variant = infer_variant(args.checkpoint, sefi_config, args.variant)
254353
transformer_config = build_transformer_config(root, sefi_config)
255354

256-
transformer = SeFiTransformer2DModel(**transformer_config)
257-
state_dict = load_transformer_state_dict(root / "transformer")
258-
missing, unexpected = transformer.load_state_dict(state_dict, strict=False)
259-
if missing or unexpected:
260-
raise ValueError(f"Transformer state dict mismatch. Missing={missing[:20]}, unexpected={unexpected[:20]}")
355+
state_dict = convert_transformer_state_dict(load_transformer_state_dict(root / "transformer"))
356+
with torch.device("meta"):
357+
transformer = SeFiTransformer2DModel(**transformer_config)
358+
transformer.load_state_dict(state_dict, strict=True, assign=True)
359+
expected_transformer_keys = set(state_dict)
360+
expected_transformer_dtypes = {key: value.dtype for key, value in state_dict.items()}
361+
floating_dtypes = {value.dtype for value in state_dict.values() if value.is_floating_point()}
362+
if len(floating_dtypes) != 1:
363+
raise ValueError(f"Expected one floating-point transformer dtype, got {sorted(map(str, floating_dtypes))}.")
364+
transformer_dtype = floating_dtypes.pop()
261365
transformer.save_pretrained(output / "transformer", safe_serialization=True)
366+
del transformer, state_dict
367+
368+
reloaded_transformer = SeFiTransformer2DModel.from_pretrained(output / "transformer", dtype=transformer_dtype)
369+
if set(reloaded_transformer.state_dict()) != expected_transformer_keys:
370+
raise ValueError("Transformer state dict keys changed after the save/load round trip.")
371+
round_trip_dtypes = {key: value.dtype for key, value in reloaded_transformer.state_dict().items()}
372+
if round_trip_dtypes != expected_transformer_dtypes:
373+
raise ValueError("Transformer state dict dtypes changed after the save/load round trip.")
374+
del reloaded_transformer
262375

263-
copytree(root / "scheduler", output / "scheduler")
376+
convert_scheduler(root, output)
264377
copytree(root / "vae", output / "vae")
265378

266379
text_encoder_name = sefi_config["model"]["text_encoder"]["model_name"]
@@ -299,6 +412,10 @@ def main():
299412
"max_sequence_length": int(model_config["text_encoder"].get("max_length", 1024)),
300413
}
301414
save_json(output / "model_index.json", model_index)
415+
source_repo_id = args.source_repo_id
416+
if source_repo_id is None and not Path(args.checkpoint).expanduser().exists():
417+
source_repo_id = args.checkpoint
418+
save_model_card(output, source_repo_id, args.target_repo_id or output.name, variant)
302419
shutil.copy2(root / "sefi_config.yaml", output / "sefi_config.yaml")
303420
print(f"Saved SeFi-Image Diffusers checkpoint to {output}")
304421

src/diffusers/models/transformers/transformer_sefi.py

Lines changed: 56 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -19,10 +19,11 @@
1919
import torch.nn as nn
2020

2121
from ...configuration_utils import ConfigMixin, register_to_config
22-
from ...utils import BaseOutput, apply_lora_scale
22+
from ...utils import BaseOutput
2323
from ..embeddings import TimestepEmbedding, Timesteps
2424
from ..modeling_utils import ModelMixin
25-
from .transformer_flux2 import Flux2Transformer2DModel
25+
from ..normalization import AdaLayerNormContinuous
26+
from .transformer_flux2 import Flux2Modulation, Flux2PosEmbed, Flux2SingleTransformerBlock, Flux2TransformerBlock
2627

2728

2829
@dataclass
@@ -105,8 +106,6 @@ class SeFiTransformer2DModel(ModelMixin, ConfigMixin):
105106
RoPE theta.
106107
eps (`float`, defaults to `1e-6`):
107108
Normalization epsilon.
108-
text_input_dim (`int`, *optional*):
109-
Expected text embedding dimension. Defaults to `joint_attention_dim`.
110109
"""
111110

112111
_supports_gradient_checkpointing = True
@@ -130,44 +129,60 @@ def __init__(
130129
axes_dims_rope: tuple[int, ...] = (32, 32, 32, 32),
131130
rope_theta: int = 2000,
132131
eps: float = 1e-6,
133-
text_input_dim: int | None = None,
134132
):
135133
super().__init__()
136134

137-
text_input_dim = joint_attention_dim if text_input_dim is None else text_input_dim
138-
if int(text_input_dim) != int(joint_attention_dim):
139-
raise ValueError(
140-
f"`text_input_dim` must match `joint_attention_dim`, got {text_input_dim} and {joint_attention_dim}."
141-
)
142-
143135
self.out_channels = out_channels or in_channels
144136
self.inner_dim = num_attention_heads * attention_head_dim
145-
self.backbone = Flux2Transformer2DModel(
146-
patch_size=patch_size,
147-
in_channels=in_channels,
148-
out_channels=out_channels,
149-
num_layers=num_layers,
150-
num_single_layers=num_single_layers,
151-
attention_head_dim=attention_head_dim,
152-
num_attention_heads=num_attention_heads,
153-
joint_attention_dim=joint_attention_dim,
154-
timestep_guidance_channels=timestep_guidance_channels,
155-
mlp_ratio=mlp_ratio,
156-
axes_dims_rope=axes_dims_rope,
157-
rope_theta=rope_theta,
158-
eps=eps,
159-
guidance_embeds=False,
160-
)
161-
# The reference SeFi transformer deletes Flux2's timestep/guidance embedder and stores only the dual embedder.
162-
self.backbone.time_guidance_embed = nn.Identity()
137+
138+
self.pos_embed = Flux2PosEmbed(theta=rope_theta, axes_dim=axes_dims_rope)
163139
self.dual_time_embed = SeFiDualTimestepEmbeddings(
164140
in_channels=timestep_guidance_channels,
165141
embedding_dim=self.inner_dim,
166142
bias=False,
167143
)
144+
145+
self.double_stream_modulation_img = Flux2Modulation(self.inner_dim, mod_param_sets=2, bias=False)
146+
self.double_stream_modulation_txt = Flux2Modulation(self.inner_dim, mod_param_sets=2, bias=False)
147+
self.single_stream_modulation = Flux2Modulation(self.inner_dim, mod_param_sets=1, bias=False)
148+
149+
self.x_embedder = nn.Linear(in_channels, self.inner_dim, bias=False)
150+
self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim, bias=False)
151+
152+
self.transformer_blocks = nn.ModuleList(
153+
[
154+
Flux2TransformerBlock(
155+
dim=self.inner_dim,
156+
num_attention_heads=num_attention_heads,
157+
attention_head_dim=attention_head_dim,
158+
mlp_ratio=mlp_ratio,
159+
eps=eps,
160+
bias=False,
161+
)
162+
for _ in range(num_layers)
163+
]
164+
)
165+
self.single_transformer_blocks = nn.ModuleList(
166+
[
167+
Flux2SingleTransformerBlock(
168+
dim=self.inner_dim,
169+
num_attention_heads=num_attention_heads,
170+
attention_head_dim=attention_head_dim,
171+
mlp_ratio=mlp_ratio,
172+
eps=eps,
173+
bias=False,
174+
)
175+
for _ in range(num_single_layers)
176+
]
177+
)
178+
179+
self.norm_out = AdaLayerNormContinuous(
180+
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=eps, bias=False
181+
)
182+
self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=False)
183+
168184
self.gradient_checkpointing = False
169185

170-
@apply_lora_scale("joint_attention_kwargs")
171186
def forward(
172187
self,
173188
hidden_states: torch.Tensor,
@@ -211,26 +226,26 @@ def forward(
211226
timestep_tex = timestep_tex.to(hidden_states.dtype) * 1000
212227
temb = self.dual_time_embed(timestep_sem, timestep_tex)
213228

214-
double_stream_mod_img = self.backbone.double_stream_modulation_img(temb)
215-
double_stream_mod_txt = self.backbone.double_stream_modulation_txt(temb)
216-
single_stream_mod = self.backbone.single_stream_modulation(temb)
229+
double_stream_mod_img = self.double_stream_modulation_img(temb)
230+
double_stream_mod_txt = self.double_stream_modulation_txt(temb)
231+
single_stream_mod = self.single_stream_modulation(temb)
217232

218-
hidden_states = self.backbone.x_embedder(hidden_states)
219-
encoder_hidden_states = self.backbone.context_embedder(encoder_hidden_states)
233+
hidden_states = self.x_embedder(hidden_states)
234+
encoder_hidden_states = self.context_embedder(encoder_hidden_states)
220235

221236
if img_ids.ndim == 3:
222237
img_ids = img_ids[0]
223238
if txt_ids.ndim == 3:
224239
txt_ids = txt_ids[0]
225240

226-
image_rotary_emb = self.backbone.pos_embed(img_ids)
227-
text_rotary_emb = self.backbone.pos_embed(txt_ids)
241+
image_rotary_emb = self.pos_embed(img_ids)
242+
text_rotary_emb = self.pos_embed(txt_ids)
228243
concat_rotary_emb = (
229244
torch.cat([text_rotary_emb[0], image_rotary_emb[0]], dim=0),
230245
torch.cat([text_rotary_emb[1], image_rotary_emb[1]], dim=0),
231246
)
232247

233-
for block in self.backbone.transformer_blocks:
248+
for block in self.transformer_blocks:
234249
if torch.is_grad_enabled() and self.gradient_checkpointing:
235250
encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
236251
block,
@@ -253,7 +268,7 @@ def forward(
253268

254269
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
255270

256-
for block in self.backbone.single_transformer_blocks:
271+
for block in self.single_transformer_blocks:
257272
if torch.is_grad_enabled() and self.gradient_checkpointing:
258273
hidden_states = self._gradient_checkpointing_func(
259274
block,
@@ -273,8 +288,8 @@ def forward(
273288
)
274289

275290
hidden_states = hidden_states[:, num_txt_tokens:, ...]
276-
hidden_states = self.backbone.norm_out(hidden_states, temb)
277-
output = self.backbone.proj_out(hidden_states)
291+
hidden_states = self.norm_out(hidden_states, temb)
292+
output = self.proj_out(hidden_states)
278293

279294
if not return_dict:
280295
return (output,)

0 commit comments

Comments
 (0)