From fb521acda78a50ad60497137bb4e7be7f705b9bc Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 8 Feb 2024 16:48:27 -0700 Subject: [PATCH] Hopefully fix an issue where BlehDeepShrink fadeout could cause tensor size mismatches --- py/deepshrink.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/py/deepshrink.py b/py/deepshrink.py index 17f2497..da27633 100644 --- a/py/deepshrink.py +++ b/py/deepshrink.py @@ -128,10 +128,13 @@ class DeepShrinkBleh: / (end_percent - start_fadeout_percent) ) scaled_scale = 1.0 - ((1.0 - downscale_factor) * downscale_pct) + orig_width, orig_height = h.shape[-1], h.shape[-2] width, height = ( - round(h.shape[-1] * scaled_scale), - round(h.shape[-2] * scaled_scale), + round(orig_width * scaled_scale), + round(orig_height * scaled_scale), ) + if scaled_scale >= 0.98 or width >= orig_width or height >= orig_height: + return h if downscale_method == "bislerp": return bislerp(h, width, height) return torch.nn.functional.interpolate(