From 890789d235b345e5391b457406859e5cb3a70a4e Mon Sep 17 00:00:00 2001 From: lindicaphxag-tech Date: Fri, 2 Oct 2026 17:19:24 +0800 Subject: [PATCH 1/4] fix(unet3d): apply class conditioning --- src/diffusers/models/unets/unet_3d_condition.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/src/diffusers/models/unets/unet_3d_condition.py b/src/diffusers/models/unets/unet_3d_condition.py index 0d15e93da68f..2ca097d8090e 100644 --- a/src/diffusers/models/unets/unet_3d_condition.py +++ b/src/diffusers/models/unets/unet_3d_condition.py @@ -90,6 +90,9 @@ class UNet3DConditionModel(ModelMixin, AttentionMixin, ConfigMixin, UNet2DCondit num_attention_heads (`int`, *optional*): The number of attention heads. time_cond_proj_dim (`int`, *optional*, defaults to `None`): The dimension of `cond_proj` layer in the timestep embedding. + num_class_embeds (`int`, *optional*, defaults to `None`): + Input dimension of an optional learnable class embedding. When set, `class_labels` are embedded and + added to the timestep embeddings. """ _supports_gradient_checkpointing = False @@ -124,6 +127,7 @@ def __init__( attention_head_dim: int | tuple[int] = 64, num_attention_heads: int | tuple[int] | None = None, time_cond_proj_dim: int | None = None, + num_class_embeds: int | None = None, ): super().__init__() @@ -177,6 +181,9 @@ def __init__( act_fn=act_fn, cond_proj_dim=time_cond_proj_dim, ) + self.class_embedding = ( + nn.Embedding(num_class_embeds, time_embed_dim) if num_class_embeds is not None else None + ) self.transformer_in = TransformerTemporalModel( num_attention_heads=8, @@ -567,6 +574,11 @@ def forward( t_emb = t_emb.to(dtype=self.dtype) emb = self.time_embedding(t_emb, timestep_cond) + if self.class_embedding is not None: + if class_labels is None: + raise ValueError("class_labels should be provided when num_class_embeds > 0") + emb = emb + self.class_embedding(class_labels).to(dtype=emb.dtype) + emb = emb.repeat_interleave(num_frames, dim=0, output_size=emb.shape[0] * num_frames) encoder_hidden_states = encoder_hidden_states.repeat_interleave( num_frames, dim=0, output_size=encoder_hidden_states.shape[0] * num_frames From ac29ecac74e02bce0fad9b3427fb2772349da0d5 Mon Sep 17 00:00:00 2001 From: lindicaphxag-tech Date: Fri, 2 Oct 2026 17:19:27 +0800 Subject: [PATCH 2/4] test(unet3d): cover class labels --- .../unets/test_models_unet_3d_condition.py | 23 +++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/tests/models/unets/test_models_unet_3d_condition.py b/tests/models/unets/test_models_unet_3d_condition.py index bb05c8d87058..97d0b6ff2b4e 100644 --- a/tests/models/unets/test_models_unet_3d_condition.py +++ b/tests/models/unets/test_models_unet_3d_condition.py @@ -91,6 +91,29 @@ def test_forward_with_norm_groups(self): assert output.shape == self.get_dummy_inputs()["sample"].shape, "Input and output shapes do not match" + def test_class_conditioning(self): + init_dict = self.get_init_dict() + init_dict["num_class_embeds"] = 2 + model = self.model_class(**init_dict).to(torch_device).eval() + + inputs = self.get_dummy_inputs() + batch_size = inputs["sample"].shape[0] + labels_0 = torch.zeros(batch_size, dtype=torch.long, device=torch_device) + labels_1 = torch.ones(batch_size, dtype=torch.long, device=torch_device) + + with torch.no_grad(): + output_0 = model(**inputs, class_labels=labels_0).sample + output_1 = model(**inputs, class_labels=labels_1).sample + + assert not torch.equal(output_0, output_1) + + try: + model(**inputs) + except ValueError as error: + assert "class_labels should be provided" in str(error) + else: + raise AssertionError("Expected class-conditioned UNet3D to require class_labels") + def test_feed_forward_chunking(self): init_dict = self.get_init_dict() init_dict["block_out_channels"] = (32, 64) From 5ca67dcfda512b731c4cdb4c9c1889902485b012 Mon Sep 17 00:00:00 2001 From: lindicaphxag-tech Date: Fri, 2 Oct 2026 18:04:51 +0800 Subject: [PATCH 3/4] style: format UNet3D class embedding --- src/diffusers/models/unets/unet_3d_condition.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/diffusers/models/unets/unet_3d_condition.py b/src/diffusers/models/unets/unet_3d_condition.py index 2ca097d8090e..e3766eb863e1 100644 --- a/src/diffusers/models/unets/unet_3d_condition.py +++ b/src/diffusers/models/unets/unet_3d_condition.py @@ -181,9 +181,7 @@ def __init__( act_fn=act_fn, cond_proj_dim=time_cond_proj_dim, ) - self.class_embedding = ( - nn.Embedding(num_class_embeds, time_embed_dim) if num_class_embeds is not None else None - ) + self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) if num_class_embeds is not None else None self.transformer_in = TransformerTemporalModel( num_attention_heads=8, From 08f3951154c359efaabc4f20642ae94605dc823e Mon Sep 17 00:00:00 2001 From: lindicaphxag-tech Date: Fri, 2 Oct 2026 18:06:51 +0800 Subject: [PATCH 4/4] docs: align UNet3D class embedding docs --- src/diffusers/models/unets/unet_3d_condition.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/diffusers/models/unets/unet_3d_condition.py b/src/diffusers/models/unets/unet_3d_condition.py index e3766eb863e1..f8038993d3de 100644 --- a/src/diffusers/models/unets/unet_3d_condition.py +++ b/src/diffusers/models/unets/unet_3d_condition.py @@ -91,8 +91,8 @@ class UNet3DConditionModel(ModelMixin, AttentionMixin, ConfigMixin, UNet2DCondit time_cond_proj_dim (`int`, *optional*, defaults to `None`): The dimension of `cond_proj` layer in the timestep embedding. num_class_embeds (`int`, *optional*, defaults to `None`): - Input dimension of an optional learnable class embedding. When set, `class_labels` are embedded and - added to the timestep embeddings. + Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing + class conditioning. """ _supports_gradient_checkpointing = False