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
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,7 @@ def test_save_load_components(self):
pipe = self.get_pipeline()

with tempfile.TemporaryDirectory() as tmpdir:
pipe.save_pretrained(tmpdir, safe_serialization=True)
pipe.save_pretrained(tmpdir, safe_serialization=True, overwrite_modular_index=True)
pipe = self.pipeline_class.from_pretrained(tmpdir)
pipe.load_components()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -514,7 +514,7 @@ def test_set_timesteps_native_flow_schedule(self):
class TestCosmos3OmniModularPipelineLoading(Cosmos3OmniModularPipelineTesterConfig, ModularLoadingTesterMixin):
def test_save_from_pretrained(self, tmp_path):
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

loaded_pipe = ModularPipeline.from_pretrained(str(tmp_path))
loaded_pipe.load_components(dtype=torch.float32)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -205,7 +205,7 @@ class TestCosmos3DistilledModularPipelineLoading(
):
def test_save_from_pretrained(self, tmp_path):
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

loaded_pipe = ModularPipeline.from_pretrained(str(tmp_path))
loaded_pipe.load_components(torch_dtype=torch.float32)
Expand Down
4 changes: 2 additions & 2 deletions tests/modular_pipelines/flux/test_modular_pipeline_flux.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,7 +162,7 @@ def test_float16_inference(self):
class TestFluxImg2ImgModularPipelineLoading(FluxImg2ImgModularPipelineTesterConfig, ModularLoadingTesterMixin):
def test_save_from_pretrained(self, tmp_path, base_pipe_output):
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = ModularPipeline.from_pretrained(tmp_path)
pipe.load_components(dtype=torch.float32)
Expand Down Expand Up @@ -249,7 +249,7 @@ def test_float16_inference(self):
class TestFluxKontextModularPipelineLoading(FluxKontextModularPipelineTesterConfig, ModularLoadingTesterMixin):
def test_save_from_pretrained(self, tmp_path, base_pipe_output):
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = ModularPipeline.from_pretrained(tmp_path)
pipe.load_components(dtype=torch.float32)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,7 @@ class TestMiniMaxMusic3ModularPipelineLoading(MiniMaxMusic3ModularPipelineTester
def test_save_from_pretrained(self, tmp_path, base_pipe_output):
# the common implementation indexes 4-D image outputs; compare the audio waveform directly
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = ModularPipeline.from_pretrained(str(tmp_path))
pipe.load_components(dtype=torch.float32)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ def test_save_from_pretrained(self, tmp_path):
base_pipe = self.get_pipeline().to(torch_device)
pipes.append(base_pipe)

base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)
pipe = self.pipeline_class.from_pretrained(tmp_path).to(torch_device)
pipe.load_components(dtype=torch.float32)
pipe.to(torch_device)
Expand All @@ -102,7 +102,7 @@ def test_save_from_pretrained(self, tmp_path):

def test_load_expected_components_from_save_pretrained(self, tmp_path):
base_pipe = self.get_pipeline()
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = self.pipeline_class.from_pretrained(tmp_path)
pipe.load_components(dtype=torch.float32)
Expand Down Expand Up @@ -191,7 +191,7 @@ def test_save_from_pretrained(self, tmp_path):
base_pipe = self.get_pipeline().to(torch_device)
pipes.append(base_pipe)

base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)
pipe = self.pipeline_class.from_pretrained(tmp_path).to(torch_device)
pipe.load_components(dtype=torch.float32)
pipe.to(torch_device)
Expand All @@ -208,7 +208,7 @@ def test_save_from_pretrained(self, tmp_path):

def test_load_expected_components_from_save_pretrained(self, tmp_path):
base_pipe = self.get_pipeline()
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = self.pipeline_class.from_pretrained(tmp_path)
pipe.load_components(dtype=torch.float32)
Expand Down
17 changes: 14 additions & 3 deletions tests/modular_pipelines/test_modular_pipeline_loading.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
import json
import os
import shutil
import warnings

import pytest
import torch
Expand Down Expand Up @@ -156,7 +157,7 @@ def test_save_pretrained_roundtrip_with_local_model(self, tmp_path):
original_state_dict = pipe.unet.state_dict()

save_dir = str(tmp_path / "my-pipeline")
pipe.save_pretrained(save_dir)
pipe.save_pretrained(save_dir, overwrite_modular_index=True)

loaded_pipe = ModularPipeline.from_pretrained(save_dir)
loaded_pipe.load_components(dtype=torch.float32)
Expand Down Expand Up @@ -204,7 +205,8 @@ def test_save_pretrained_default_writes_self_contained_local_copy(self, tmp_path
pipe.load_components(names=["unet"], dtype=torch.float32)

save_dir = str(tmp_path / "my-pipeline")
pipe.save_pretrained(save_dir)
with pytest.warns(FutureWarning, match="overwrite_modular_index"):
pipe.save_pretrained(save_dir)

with open(os.path.join(save_dir, "modular_model_index.json")) as f:
index = json.load(f)
Expand All @@ -221,6 +223,15 @@ def test_save_pretrained_default_writes_self_contained_local_copy(self, tmp_path
loaded_pipe.load_components(names=["unet"], dtype=torch.float32, local_files_only=True)
assert loaded_pipe.unet is not None

@pytest.mark.parametrize("overwrite_modular_index", [True, False])
def test_save_pretrained_explicit_overwrite_modular_index_does_not_warn(self, tmp_path, overwrite_modular_index):
pipe = ModularPipeline.from_pretrained("hf-internal-testing/tiny-stable-diffusion-xl-pipe")
pipe.load_components(names=["unet"], dtype=torch.float32)

with warnings.catch_warnings():
warnings.filterwarnings("error", message=".*overwrite_modular_index", category=FutureWarning)
pipe.save_pretrained(str(tmp_path / "my-pipeline"), overwrite_modular_index=overwrite_modular_index)

def test_save_pretrained_overwrite_modular_index(self, tmp_path):
"""With overwrite_modular_index=True, all component references should point to the save directory."""
pipe = ModularPipeline.from_pretrained("hf-internal-testing/tiny-stable-diffusion-xl-pipe")
Expand Down Expand Up @@ -266,7 +277,7 @@ def test_init_fallback_when_blocks_class_name_is_base_class(self, tmp_path):

# 3. Save and reload — the saved config will have _blocks_class_name="SequentialPipelineBlocks"
save_dir = str(tmp_path / "pipeline")
t2i_pipe.save_pretrained(save_dir)
t2i_pipe.save_pretrained(save_dir, overwrite_modular_index=True)
loaded_pipe = ModularPipeline.from_pretrained(save_dir)

# 4. Verify it fell back to default_blocks_name and has correct blocks
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -364,7 +364,7 @@ def __call__(self, components, state: PipelineState) -> PipelineState:

block = ModularPipelineBlocks.from_pretrained(pipeline_repo_dir, trust_remote_code=True)
pipe = block.init_pipeline()
pipe.save_pretrained(pipeline_repo_dir)
pipe.save_pretrained(pipeline_repo_dir, overwrite_modular_index=True)

# Step 3: Load the pipeline from the saved directory.
loaded_pipe = ModularPipeline.from_pretrained(pipeline_repo_dir, trust_remote_code=True)
Expand Down
4 changes: 2 additions & 2 deletions tests/modular_pipelines/testing_utils/loading.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ class ModularLoadingTesterMixin(BaseModularPipelineOutputMixin):

def test_save_from_pretrained(self, tmp_path, base_pipe_output):
base_pipe = self.get_pipeline().to(torch_device)
base_pipe.save_pretrained(str(tmp_path))
base_pipe.save_pretrained(str(tmp_path), overwrite_modular_index=True)

pipe = ModularPipeline.from_pretrained(tmp_path)
pipe.load_components(dtype=torch.float32)
Expand Down Expand Up @@ -67,7 +67,7 @@ def test_load_expected_components_from_pretrained(self, tmp_path):
def test_load_expected_components_from_save_pretrained(self, tmp_path):
pipe = self.get_pipeline()
save_dir = str(tmp_path / "saved-pipeline")
pipe.save_pretrained(save_dir)
pipe.save_pretrained(save_dir, overwrite_modular_index=True)

expected = get_specified_components(save_dir)
loaded_pipe = ModularPipeline.from_pretrained(save_dir)
Expand Down
Loading