From 293aab8a411e4ca1010287ce870a363c928743ef Mon Sep 17 00:00:00 2001 From: Devam0311 Date: Tue, 1 Sep 2026 15:52:09 +0530 Subject: [PATCH] Do not mutate the source transformer config in SD3ControlNetModel.from_transformer `from_transformer` bound `config = transformer.config`, aliasing the transformer's live config, then wrote `num_layers` and `extra_conditioning_channels` into it. Building a ControlNet therefore silently changed the source transformer, which can corrupt later serialization, logging, or pipeline construction that reuses it. Copy the config first, as FluxControlNetModel and QwenImageControlNetModel already do. Add regression tests covering the config staying untouched, the ControlNet receiving the overrides, and `num_layers=None` falling back to the transformer's value. Ref #13611 (Issue 4) Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01XBYeq5vB4DNZDEaqUsroKR --- .../models/controlnets/controlnet_sd3.py | 4 +- .../controlnets/test_models_controlnet_sd3.py | 83 +++++++++++++++++++ 2 files changed, 85 insertions(+), 2 deletions(-) create mode 100644 tests/models/controlnets/test_models_controlnet_sd3.py diff --git a/src/diffusers/models/controlnets/controlnet_sd3.py b/src/diffusers/models/controlnets/controlnet_sd3.py index 1f0ca529ff16..2ac1b66473b1 100644 --- a/src/diffusers/models/controlnets/controlnet_sd3.py +++ b/src/diffusers/models/controlnets/controlnet_sd3.py @@ -254,8 +254,8 @@ def _get_pos_embed_from_transformer(self, transformer): def from_transformer( cls, transformer, num_layers=12, num_extra_conditioning_channels=1, load_weights_from_transformer=True ): - config = transformer.config - config["num_layers"] = num_layers or config.num_layers + config = dict(transformer.config) + config["num_layers"] = num_layers or transformer.config.num_layers config["extra_conditioning_channels"] = num_extra_conditioning_channels controlnet = cls.from_config(config) diff --git a/tests/models/controlnets/test_models_controlnet_sd3.py b/tests/models/controlnets/test_models_controlnet_sd3.py new file mode 100644 index 000000000000..854cdfc22d97 --- /dev/null +++ b/tests/models/controlnets/test_models_controlnet_sd3.py @@ -0,0 +1,83 @@ +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch + +from diffusers import SD3ControlNetModel, SD3Transformer2DModel + +from ...testing_utils import enable_full_determinism + + +enable_full_determinism() + + +def get_dummy_transformer(): + torch.manual_seed(0) + return SD3Transformer2DModel( + sample_size=4, + patch_size=1, + in_channels=4, + out_channels=4, + num_layers=3, + attention_head_dim=4, + num_attention_heads=2, + caption_projection_dim=8, + joint_attention_dim=8, + pooled_projection_dim=8, + ) + + +class TestSD3ControlNetModelFromTransformer: + def test_from_transformer_does_not_mutate_source_config(self): + # Regression: `from_transformer` aliased the transformer's live config and wrote + # ControlNet-specific values into it, so building a ControlNet silently changed the + # source transformer's `num_layers` and added `extra_conditioning_channels`. + transformer = get_dummy_transformer() + config_before = dict(transformer.config) + + SD3ControlNetModel.from_transformer( + transformer, + num_layers=1, + num_extra_conditioning_channels=2, + load_weights_from_transformer=False, + ) + + assert dict(transformer.config) == config_before, ( + "`from_transformer` must not modify the source transformer's config." + ) + + def test_from_transformer_applies_controlnet_config(self): + transformer = get_dummy_transformer() + + controlnet = SD3ControlNetModel.from_transformer( + transformer, + num_layers=1, + num_extra_conditioning_channels=2, + load_weights_from_transformer=False, + ) + + assert controlnet.config.num_layers == 1 + assert controlnet.config.extra_conditioning_channels == 2 + + def test_from_transformer_num_layers_falls_back_to_transformer(self): + transformer = get_dummy_transformer() + + controlnet = SD3ControlNetModel.from_transformer( + transformer, + num_layers=None, + num_extra_conditioning_channels=1, + load_weights_from_transformer=False, + ) + + assert controlnet.config.num_layers == transformer.config.num_layers