diff --git a/adv_control/utils.py b/adv_control/utils.py index 2d380ed..1ace23a 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -359,7 +359,7 @@ def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim mask = mask.clone() if flux_shape is not None: multiplier = multiplier * 0.5 - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(flux_shape[-2]*multiplier), round(flux_shape[-1]*multiplier)), mode="bilinear") + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(math.ceil(flux_shape[-2]*multiplier), math.ceil(flux_shape[-1]*multiplier)), mode="bilinear") mask = rearrange(mask, "b c h w -> b (h w) c") else: mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(shape[-2]*multiplier), round(shape[-1]*multiplier)), mode="bilinear") diff --git a/tests/test_modern_control.py b/tests/test_modern_control.py index b4b173a..7ca4f70 100644 --- a/tests/test_modern_control.py +++ b/tests/test_modern_control.py @@ -63,6 +63,20 @@ class ModernControlPreprocessingTests(unittest.TestCase): ) torch.testing.assert_close(output, expected) + def test_effect_mask_matches_padded_flux_tokens_for_odd_latent_size(self): + control = ControlNetAdvanced(ControlModel(), None) + control.x_noisy_shape = (1, 16, 5, 7) + control.mask_cond_hint = torch.ones((1, 1, 5, 7)) + control.tk_mask_cond_hint = None + control.weights = SimpleNamespace(has_uncond_multiplier=False, has_uncond_mask=False) + control.latent_keyframes = None + control._current_timestep_keyframe = SimpleNamespace(strength=1.0) + + output = torch.ones((1, 12, 4)) + control.apply_advanced_strengths_and_masks(output, batched_number=1) + + torch.testing.assert_close(output, torch.ones_like(output)) + def test_vae_compression_and_source_mask_match_5d_hint(self): control_model = ControlModel() vae = VideoVAE()