Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions .ai/skills/custom-blocks/SKILL.md
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,8 @@ diffusers-cli custom_blocks [--block_module_name <file.py>] [--block_class_name
### What it does

1. **AST scan**: parses `<file>` 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
Expand Down Expand Up @@ -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 <file>: [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 `<ClassName>()` 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__`.
Expand Down
5 changes: 3 additions & 2 deletions docs/source/en/using-diffusers/cli.md
Original file line number Diff line number Diff line change
Expand Up @@ -290,8 +290,9 @@ hf sandbox kill <id>
## `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
Expand Down
14 changes: 12 additions & 2 deletions src/diffusers/commands/custom_blocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
28 changes: 27 additions & 1 deletion tests/others/test_cli_commands.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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):
Expand Down
Loading