From 61fa161c34c2ecf4cd9a5d4447c08bda3b53585a Mon Sep 17 00:00:00 2001 From: Acly Date: Thu, 12 Sep 2024 10:17:37 +0200 Subject: [PATCH] Regions: also check if dtype matches --- region.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/region.py b/region.py index 1aab35d..292bd28 100644 --- a/region.py +++ b/region.py @@ -150,9 +150,9 @@ class AttentionMask: assert k.mean() == v.mean(), "k and v must be the same." device, dtype = q.device, q.dtype - if self.conds[0].device != device: + if self.conds[0].device != device or self.conds[0].dtype != dtype: self.conds = [cond.to(device, dtype=dtype) for cond in self.conds] - if self.mask.device != device: + if self.mask.device != device or self.mask.dtype != dtype: self.mask = self.mask.to(device, dtype=dtype) cond_or_unconds = extra_options["cond_or_uncond"]