Skip to content

Commit 96c84cf

Browse files
docs: revert unet forward
1 parent 16f4df0 commit 96c84cf

1 file changed

Lines changed: 1 addition & 2 deletions

File tree

src/optimum/rbln/diffusers/models/unets/unet_spatio_temporal_condition.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -166,7 +166,7 @@ def forward(
166166
encoder_hidden_states: torch.Tensor,
167167
added_time_ids: torch.Tensor,
168168
return_dict: bool = True,
169-
**kwargs: Any,
169+
**kwargs,
170170
) -> Union[UNetSpatioTemporalConditionOutput, Tuple]:
171171
"""
172172
Forward pass for the RBLN-optimized UNet2DConditionModel.
@@ -176,7 +176,6 @@ def forward(
176176
encoder_hidden_states (torch.Tensor): The encoder hidden states.
177177
added_time_ids (torch.Tensor): A tensor containing additional sinusoidal embeddings and added to the time embeddings.
178178
return_dict (bool): Whether or not to return a [`~diffusers.models.unets.unet_spatio_temporal.UNetSpatioTemporalConditionOutput`] instead of a plain tuple.
179-
kwargs: Additional arguments for the forward method.
180179
181180
Returns:
182181
(Union[`~diffusers.models.unets.unet_spatio_temporal.UNetSpatioTemporalConditionOutput`], Tuple)

0 commit comments

Comments
 (0)