From 47b5fd5e3d76e4deeb334f0fdabf063bfb308688 Mon Sep 17 00:00:00 2001 From: yiyixuxu Date: Fri, 9 Oct 2026 21:42:47 +0200 Subject: [PATCH] [cli] detect composite block classes in `custom_blocks` `diffusers-cli custom_blocks` only recognized classes whose base was spelled exactly `ModularPipelineBlocks`, so a `SequentialPipelineBlocks` or `AutoPipelineBlocks` subclass (the usual thing to publish) was reported as "could not be retrieved". Accept the composite subclasses as well, and match on the class name so dotted bases like `diffusers.modular_pipelines.SequentialPipelineBlocks` work too. Co-Authored-By: Claude Fable 5.1 --- .ai/skills/custom-blocks/SKILL.md | 9 ++++---- docs/source/en/using-diffusers/cli.md | 5 +++-- src/diffusers/commands/custom_blocks.py | 14 +++++++++++-- tests/others/test_cli_commands.py | 28 ++++++++++++++++++++++++- 4 files changed, 47 insertions(+), 9 deletions(-) diff --git a/.ai/skills/custom-blocks/SKILL.md b/.ai/skills/custom-blocks/SKILL.md index b98aa18b617e..261b4cdcd877 100644 --- a/.ai/skills/custom-blocks/SKILL.md +++ b/.ai/skills/custom-blocks/SKILL.md @@ -57,7 +57,8 @@ diffusers-cli custom_blocks [--block_module_name ] [--block_class_name ### What it does 1. **AST scan**: parses `` without executing it, walks top-level `ClassDef` nodes, and collects every - class whose `bases` include `ModularPipelineBlocks`. + class whose `bases` include `ModularPipelineBlocks` or one of its composite subclasses + (`SequentialPipelineBlocks`, `AutoPipelineBlocks`, `ConditionalPipelineBlocks`, `LoopSequentialPipelineBlocks`). 2. **Pick a class**: uses `--block_class_name` if given, else the first found. Errors with the list of available classes if your name doesn't match. 3. **Load and save**: imports the file via `importlib.util.spec_from_file_location` (this does execute the @@ -137,9 +138,9 @@ diffusers-cli run --model my-user/my-denoise-block --trust-remote-code \ - **`block_class_name could not be retrieved. Available classes from : [ClassA, ClassB]`** — your `--block_class_name` doesn't match any `ModularPipelineBlocks` subclass found. Pick from the list shown. - **No classes found**: silent — the command will try to use the first entry in an empty list and raise - `IndexError`. If you hit that, double-check your class actually inherits from `ModularPipelineBlocks` - (the AST scan looks for that literal base-class name; aliased imports like `from diffusers import ... - as MPB` won't be picked up). + `IndexError`. If you hit that, double-check your class actually inherits from `ModularPipelineBlocks` or one + of the composite classes listed above (the AST scan looks for those literal class names; aliased imports like + `from diffusers import ... as MPB` won't be picked up). - **Block requires constructor args**: the command calls `()` with no args. If your block needs `__init__` parameters, refactor to take them from `state`/`components` at `__call__` time instead, or hardcode defaults in `__init__`. diff --git a/docs/source/en/using-diffusers/cli.md b/docs/source/en/using-diffusers/cli.md index 01339f4628a2..8625001cef89 100644 --- a/docs/source/en/using-diffusers/cli.md +++ b/docs/source/en/using-diffusers/cli.md @@ -290,8 +290,9 @@ hf sandbox kill ## `custom_blocks` Package a local `ModularPipelineBlocks` subclass for upload to the Hub. Reads a Python file, AST-scans it for -subclasses of `ModularPipelineBlocks`, instantiates the chosen one, and calls `save_pretrained` in the current -working directory. +classes deriving from `ModularPipelineBlocks` or one of its composite subclasses (`SequentialPipelineBlocks`, +`AutoPipelineBlocks`, `ConditionalPipelineBlocks`, `LoopSequentialPipelineBlocks`), instantiates the chosen one, and +calls `save_pretrained` in the current working directory. ```bash # Package the first block found in ./block.py diff --git a/src/diffusers/commands/custom_blocks.py b/src/diffusers/commands/custom_blocks.py index a3649117e002..bc6aa257cc36 100644 --- a/src/diffusers/commands/custom_blocks.py +++ b/src/diffusers/commands/custom_blocks.py @@ -28,7 +28,14 @@ from . import BaseDiffusersCLICommand -EXPECTED_PARENT_CLASSES = ["ModularPipelineBlocks"] +# the block base class and its composite subclasses; a class deriving from any of them is a packageable block +EXPECTED_PARENT_CLASSES = [ + "ModularPipelineBlocks", + "SequentialPipelineBlocks", + "AutoPipelineBlocks", + "ConditionalPipelineBlocks", + "LoopSequentialPipelineBlocks", +] def conversion_command_factory(args: Namespace): @@ -128,7 +135,10 @@ def _get_class_names(self, file_path): if not isinstance(node, ast.ClassDef): continue - base_names = [bname for b in node.bases if (bname := self._get_base_name(b)) is not None] + # compare on the class name only, so `diffusers.modular_pipelines.SequentialPipelineBlocks` matches too + base_names = [ + bname.rsplit(".", 1)[-1] for b in node.bases if (bname := self._get_base_name(b)) is not None + ] for allowed in EXPECTED_PARENT_CLASSES: if allowed in base_names: diff --git a/tests/others/test_cli_commands.py b/tests/others/test_cli_commands.py index 9b731416df79..ba30cfb4a92b 100644 --- a/tests/others/test_cli_commands.py +++ b/tests/others/test_cli_commands.py @@ -432,9 +432,16 @@ def test_class_discovery(self, tmp_path): "class OtherBase:\n pass\n" "class NotABlock(OtherBase):\n pass\n" "class MyBlock(ModularPipelineBlocks):\n pass\n" + # composite blocks are packageable too, and the base may be referenced through its module + "class MyBlocks(SequentialPipelineBlocks):\n pass\n" + "class MyAutoBlocks(diffusers.modular_pipelines.AutoPipelineBlocks):\n pass\n" ) cmd = CustomBlocksCommand() - assert cmd._get_class_names(block_py) == [("MyBlock", "ModularPipelineBlocks")] + assert cmd._get_class_names(block_py) == [ + ("MyBlock", "ModularPipelineBlocks"), + ("MyBlocks", "SequentialPipelineBlocks"), + ("MyAutoBlocks", "AutoPipelineBlocks"), + ] broken = tmp_path / "broken.py" broken.write_text("class Broken(:\n pass\n") @@ -457,6 +464,25 @@ def test_packaging_writes_pipeline_index(self, tmp_path, monkeypatch): assert (tmp_path / "modular_config.json").exists() assert (tmp_path / "modular_model_index.json").exists() + def test_packaging_composite_block(self, tmp_path, monkeypatch): + # a `SequentialPipelineBlocks` subclass is the usual thing to publish; it must be selectable by name + block_py = tmp_path / "block.py" + block_py.write_text( + "from diffusers.modular_pipelines import ModularPipelineBlocks, SequentialPipelineBlocks\n" + "\n" + "class MyStep(ModularPipelineBlocks):\n" + " model_name = 'test'\n" + "\n" + "class MyBlocks(SequentialPipelineBlocks):\n" + " model_name = 'test'\n" + " block_classes = [MyStep]\n" + " block_names = ['step']\n" + ) + monkeypatch.chdir(tmp_path) + CustomBlocksCommand(str(block_py), "MyBlocks").run() + config = json.loads((tmp_path / "modular_config.json").read_text()) + assert config["auto_map"] == {"ModularPipelineBlocks": "block.MyBlocks"} + class TestCli: def test_toplevel_help_lists_all_commands(self):