Skip to content

Commit 29a59fd

Browse files
ElectricGoaldg845
andauthored
Fix DreamLite legacy block type aliases (#14066)
* Fix DreamLite legacy block type aliases * Cover DreamLite non-V1 up block alias --------- Co-authored-by: dg845 <58458699+dg845@users.noreply.github.com>
1 parent 1ec68b6 commit 29a59fd

2 files changed

Lines changed: 99 additions & 3 deletions

File tree

src/diffusers/models/unets/unet_dreamlite.py

Lines changed: 26 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1340,10 +1340,27 @@ def forward(
13401340
# Local block dispatch (DreamLite-only)
13411341
#
13421342
# The string ``down_block_type`` / ``up_block_type`` / ``mid_block_type`` keys
1343-
# persisted in saved checkpoints' ``config.json`` mirror the Python class names
1344-
# defined above. The ``carlofkl/DreamLite-{base,mobile}`` Hub repos
1345-
# (``diffusers`` branch) ship configs that use these exact keys.
1343+
# persisted in saved checkpoints' ``config.json`` usually mirror the Python class
1344+
# names defined above. Some configs use upstream UNet block names instead.
13461345
# ---------------------------------------------------------------------------
1346+
_DREAMLITE_DOWN_BLOCK_ALIASES = {
1347+
"CrossAttnDownRemoveSelfAttnBlock2D": "DreamLiteCrossAttnNoSelfAttnDownBlock2D",
1348+
"CrossAttnDownBlock2D": "DreamLiteCrossAttnDownBlock2D",
1349+
"DownBlock2D": "DreamLiteDownBlock2D",
1350+
}
1351+
1352+
_DREAMLITE_MID_BLOCK_ALIASES = {
1353+
"UNetMidBlock2DCrossAttn": "DreamLiteUNetMidBlock2DCrossAttn",
1354+
}
1355+
1356+
_DREAMLITE_UP_BLOCK_ALIASES = {
1357+
"CrossAttnUpRemoveSelfAttnBlock2D": "DreamLiteCrossAttnNoSelfAttnUpBlock2D",
1358+
"CrossAttnUpRemoveSelfAttnBlock2DV1": "DreamLiteCrossAttnNoSelfAttnUpBlock2D",
1359+
"CrossAttnUpBlock2D": "DreamLiteCrossAttnUpBlock2D",
1360+
"UpBlock2D": "DreamLiteUpBlock2D",
1361+
}
1362+
1363+
13471364
def _get_down_block_dreamlite(
13481365
down_block_type: str,
13491366
*,
@@ -1371,6 +1388,8 @@ def _get_down_block_dreamlite(
13711388
ff_mult,
13721389
num_kv_heads,
13731390
):
1391+
down_block_type = _DREAMLITE_DOWN_BLOCK_ALIASES.get(down_block_type, down_block_type)
1392+
13741393
if down_block_type == "DreamLiteDownBlock2D":
13751394
return DreamLiteDownBlock2D(
13761395
num_layers=num_layers,
@@ -1447,6 +1466,8 @@ def _get_mid_block_dreamlite(
14471466
):
14481467
if mid_block_type is None:
14491468
return None
1469+
mid_block_type = _DREAMLITE_MID_BLOCK_ALIASES.get(mid_block_type, mid_block_type)
1470+
14501471
if mid_block_type == "DreamLiteUNetMidBlock2DCrossAttn":
14511472
return DreamLiteUNetMidBlock2DCrossAttn(
14521473
transformer_layers_per_block=transformer_layers_per_block,
@@ -1501,6 +1522,8 @@ def _get_up_block_dreamlite(
15011522
ff_mult,
15021523
num_kv_heads,
15031524
):
1525+
up_block_type = _DREAMLITE_UP_BLOCK_ALIASES.get(up_block_type, up_block_type)
1526+
15041527
if up_block_type == "DreamLiteUpBlock2D":
15051528
return DreamLiteUpBlock2D(
15061529
num_layers=num_layers,

tests/pipelines/dreamlite/test_pipeline_dreamlite.py

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,79 @@ class DreamLitePipelineFastTests(PipelineTesterMixin, unittest.TestCase):
131131
test_layerwise_casting = False
132132
test_group_offloading = False
133133

134+
def test_legacy_block_type_aliases(self):
135+
unet = DreamLiteUNetModel(
136+
sample_size=8,
137+
in_channels=4,
138+
out_channels=4,
139+
down_block_types=(
140+
"CrossAttnDownRemoveSelfAttnBlock2D",
141+
"CrossAttnDownRemoveSelfAttnBlock2D",
142+
"CrossAttnDownBlock2D",
143+
),
144+
mid_block_type="UNetMidBlock2DCrossAttn",
145+
up_block_types=(
146+
"CrossAttnUpBlock2D",
147+
"CrossAttnUpRemoveSelfAttnBlock2DV1",
148+
"UpBlock2D",
149+
),
150+
block_out_channels=(16, 32, 64),
151+
cross_attention_dim=_CROSS_ATTN_DIM,
152+
attention_head_dim=8,
153+
layers_per_block=1,
154+
norm_num_groups=8,
155+
transformer_layers_per_block=1,
156+
)
157+
158+
self.assertEqual(
159+
[block.__class__.__name__ for block in unet.down_blocks],
160+
[
161+
"DreamLiteCrossAttnNoSelfAttnDownBlock2D",
162+
"DreamLiteCrossAttnNoSelfAttnDownBlock2D",
163+
"DreamLiteCrossAttnDownBlock2D",
164+
],
165+
)
166+
self.assertEqual(unet.mid_block.__class__.__name__, "DreamLiteUNetMidBlock2DCrossAttn")
167+
self.assertEqual(
168+
[block.__class__.__name__ for block in unet.up_blocks],
169+
[
170+
"DreamLiteCrossAttnUpBlock2D",
171+
"DreamLiteCrossAttnNoSelfAttnUpBlock2D",
172+
"DreamLiteUpBlock2D",
173+
],
174+
)
175+
176+
unet_with_non_v1_up_alias = DreamLiteUNetModel(
177+
sample_size=8,
178+
in_channels=4,
179+
out_channels=4,
180+
down_block_types=(
181+
"CrossAttnDownRemoveSelfAttnBlock2D",
182+
"CrossAttnDownRemoveSelfAttnBlock2D",
183+
"CrossAttnDownBlock2D",
184+
),
185+
mid_block_type="UNetMidBlock2DCrossAttn",
186+
up_block_types=(
187+
"CrossAttnUpBlock2D",
188+
"CrossAttnUpRemoveSelfAttnBlock2D",
189+
"UpBlock2D",
190+
),
191+
block_out_channels=(16, 32, 64),
192+
cross_attention_dim=_CROSS_ATTN_DIM,
193+
attention_head_dim=8,
194+
layers_per_block=1,
195+
norm_num_groups=8,
196+
transformer_layers_per_block=1,
197+
)
198+
self.assertEqual(
199+
[block.__class__.__name__ for block in unet_with_non_v1_up_alias.up_blocks],
200+
[
201+
"DreamLiteCrossAttnUpBlock2D",
202+
"DreamLiteCrossAttnNoSelfAttnUpBlock2D",
203+
"DreamLiteUpBlock2D",
204+
],
205+
)
206+
134207
def get_dummy_components(self):
135208
torch.manual_seed(0)
136209
unet = DreamLiteUNetModel(

0 commit comments

Comments
 (0)