2323from huggingface_hub import snapshot_download
2424from safetensors .torch import load_file
2525
26- from diffusers import __version__
26+ from diffusers import FlowMatchEulerDiscreteScheduler , __version__
2727from 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+
147217def 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
167237def 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+
228327def 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
0 commit comments