mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Diffusion export bug fixed for model_index.json (#901)
## What does this PR do? **Type of change:** Bug fix <!-- Use one of the following: Bug fix, new feature, new example, new tests, documentation. --> **Overview:** Updated diffusers export to preserve the original model_index.json instead of always rebuilding a minimal one. The export now uses a simple fallback order: copy original model_index.json from source path if available, otherwise call `pipe.save_config(export_dir)`, and only then generate a minimal model_index.json as last resort. Non-diffusers export behavior is unchanged. ## Usage <!-- You can potentially add a usage example below. --> ```python # Add a code snippet demonstrating how to use this ``` ## Testing <!-- Mention how have you tested your change if applicable. --> ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes/No <!--- If No, explain why. --> - **Did you write any new necessary tests?**: Yes/No - **Did you add or update any necessary documentation?**: Yes/No - **Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**: Yes/No <!--- Only for new features, API changes, critical bug fixes or bw breaking changes. --> ## Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **Bug Fixes** * Improved Diffusers pipeline export with enhanced model configuration handling. The export process now better preserves original pipeline configurations and uses fallback strategies to ensure complete configuration files are generated. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
This commit is contained in:
@@ -937,21 +937,38 @@ def _export_diffusers_checkpoint(
|
||||
|
||||
print(f" Saved to: {component_export_dir}")
|
||||
|
||||
# Step 5: For pipelines, also save the model_index.json
|
||||
# Step 5: For pipelines, also save model_index.json
|
||||
if is_diffusers_pipe:
|
||||
model_index_path = export_dir / "model_index.json"
|
||||
if hasattr(pipe, "config") and pipe.config is not None:
|
||||
# Save a simplified model_index.json that points to the exported components
|
||||
is_partial_export = components is not None
|
||||
|
||||
# For full export, preserve original model_index.json when possible.
|
||||
# For partial export, skip this to avoid listing non-exported components.
|
||||
if not is_partial_export:
|
||||
source_path = getattr(pipe, "name_or_path", None) or getattr(
|
||||
getattr(pipe, "config", None), "_name_or_path", None
|
||||
)
|
||||
if source_path:
|
||||
candidate_model_index = Path(source_path) / "model_index.json"
|
||||
if candidate_model_index.exists():
|
||||
with open(candidate_model_index) as file:
|
||||
model_index = json.load(file)
|
||||
with open(model_index_path, "w") as file:
|
||||
json.dump(model_index, file, indent=4)
|
||||
|
||||
# Full-export fallback to Diffusers-native config serialization.
|
||||
# Partial export skips this for the same reason as above.
|
||||
if not is_partial_export and not model_index_path.exists() and hasattr(pipe, "save_config"):
|
||||
pipe.save_config(export_dir)
|
||||
|
||||
# Last resort: synthesize a minimal model_index.json from exported components.
|
||||
if not model_index_path.exists() and hasattr(pipe, "config") and pipe.config is not None:
|
||||
model_index = {
|
||||
"_class_name": type(pipe).__name__,
|
||||
"_diffusers_version": diffusers.__version__,
|
||||
}
|
||||
# Add component class names for all components
|
||||
# Use the base library name (e.g., "diffusers", "transformers") instead of
|
||||
# the full module path, as expected by diffusers pipeline loading
|
||||
for name, comp in all_components.items():
|
||||
module = type(comp).__module__
|
||||
# Extract base library name (first part of module path)
|
||||
library = module.split(".")[0]
|
||||
model_index[name] = [library, type(comp).__name__]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user