diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index e974fc6..1271abe 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -1,5 +1,4 @@ import torch -import torch from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable, get_f_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io diff --git a/more_math/Parser/MathExprParser.py b/more_math/Parser/MathExprParser.py index d5915a3..d8601f5 100644 --- a/more_math/Parser/MathExprParser.py +++ b/more_math/Parser/MathExprParser.py @@ -880,7 +880,7 @@ class MathExprParser ( Parser ): self.stmt() pass - + self.state = 69 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,1,self._ctx) @@ -917,7 +917,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcDef - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -1181,7 +1181,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_stmt - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -1920,7 +1920,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_ternaryExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -1991,7 +1991,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_compExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2251,7 +2251,7 @@ class MathExprParser ( Parser ): self.addExpr(0) pass - + self.state = 211 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,14,self._ctx) @@ -2276,7 +2276,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_addExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2396,7 +2396,7 @@ class MathExprParser ( Parser ): self.mulExpr(0) pass - + self.state = 225 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,16,self._ctx) @@ -2421,7 +2421,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_mulExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2576,7 +2576,7 @@ class MathExprParser ( Parser ): self.shiftExpr(0) pass - + self.state = 242 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,18,self._ctx) @@ -2601,7 +2601,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_shiftExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2721,7 +2721,7 @@ class MathExprParser ( Parser ): self.powExpr() pass - + self.state = 256 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,20,self._ctx) @@ -2746,7 +2746,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_powExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2839,7 +2839,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_unaryExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2954,7 +2954,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_indexExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3082,7 +3082,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_atom - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3670,7 +3670,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func0 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3730,7 +3730,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func1 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -5934,7 +5934,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func2 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -7345,7 +7345,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func3 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -7944,7 +7944,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func4 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -8147,7 +8147,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func5 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -8236,7 +8236,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcN - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -8644,7 +8644,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcNoise - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -9502,63 +9502,63 @@ class MathExprParser ( Parser ): def compExpr_sempred(self, localctx:CompExprContext, predIndex:int): if predIndex == 0: return self.precpred(self._ctx, 7) - + if predIndex == 1: return self.precpred(self._ctx, 6) - + if predIndex == 2: return self.precpred(self._ctx, 5) - + if predIndex == 3: return self.precpred(self._ctx, 4) - + if predIndex == 4: return self.precpred(self._ctx, 3) - + if predIndex == 5: return self.precpred(self._ctx, 2) - + def addExpr_sempred(self, localctx:AddExprContext, predIndex:int): if predIndex == 6: return self.precpred(self._ctx, 3) - + if predIndex == 7: return self.precpred(self._ctx, 2) - + def mulExpr_sempred(self, localctx:MulExprContext, predIndex:int): if predIndex == 8: return self.precpred(self._ctx, 4) - + if predIndex == 9: return self.precpred(self._ctx, 3) - + if predIndex == 10: return self.precpred(self._ctx, 2) - + def shiftExpr_sempred(self, localctx:ShiftExprContext, predIndex:int): if predIndex == 11: return self.precpred(self._ctx, 3) - + if predIndex == 12: return self.precpred(self._ctx, 2) - + def indexExpr_sempred(self, localctx:IndexExprContext, predIndex:int): if predIndex == 13: return self.precpred(self._ctx, 2) - + diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 7512456..8e7d267 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -923,7 +923,7 @@ class UnifiedMathVisitor(MathExprVisitor): def visitReshapeFunc(self, ctx): tsr = self._promote_to_tensor((yield ctx.expr(0))) new_shape = (yield ctx.expr(1)) - + # Ensure new_shape is a list of integers if isinstance(new_shape, torch.Tensor): new_shape = new_shape.flatten().long().tolist() @@ -941,7 +941,7 @@ class UnifiedMathVisitor(MathExprVisitor): new_shape = result elif isinstance(new_shape, (int, float)): new_shape = [int(float(new_shape))] - + return tsr.reshape(*new_shape) def visitPrintShapeFunc(self, ctx): @@ -1860,9 +1860,17 @@ class UnifiedMathVisitor(MathExprVisitor): shape_arg = self.shape if len(ctx.expr()) > 3: shape_arg = (yield ctx.expr(3)) - generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Gamma(shape_param, 1.0 / scale) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + + # Use torch.distributions.Gamma which internally handles the generator properly via torch.manual_seed + # We set the random state temporarily + old_state = torch.get_rng_state() + try: + torch.manual_seed(seed) + dist = torch.distributions.Gamma(shape_param, 1.0 / scale) + result = dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])) + return result.to(device=self.device) + finally: + torch.set_rng_state(old_state) def visitBetaDistFunc(self, ctx): seed_val = yield ctx.expr(0) @@ -1874,40 +1882,47 @@ class UnifiedMathVisitor(MathExprVisitor): shape_arg = self.shape if len(ctx.expr()) > 3: shape_arg = (yield ctx.expr(3)) - generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Beta(alpha, beta) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + + old_state = torch.get_rng_state() + try: + torch.manual_seed(seed) + dist = torch.distributions.Beta(alpha, beta) + result = dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])) + return result.to(device=self.device) + finally: + torch.set_rng_state(old_state) def visitLaplaceDistFunc(self, ctx): seed_val = yield ctx.expr(0) + shape_arg = self.shape; + if len(ctx.expr()) > 3: + shape_arg = (yield ctx.expr(3)) seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val) loc_val = yield ctx.expr(1) loc = float(loc_val.item()) if self._is_tensor(loc_val) else float(loc_val) scale_val = yield ctx.expr(2) scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val) - shape_arg = self.shape - if len(ctx.expr()) > 3: - shape_arg = (yield ctx.expr(3)) generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Laplace(loc, scale) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + return loc - scale * torch.sign(torch.empty(shape_arg, device=self.device).uniform_(-1, 1, generator=generator)) * torch.log(torch.empty(shape_arg, device=self.device).uniform_(0, 1, generator=generator).clamp(min=1e-10)) def visitGumbelDistFunc(self, ctx): seed_val = yield ctx.expr(0) + shape_arg = self.shape; + if len(ctx.expr()) > 3: + shape_arg = (yield ctx.expr(3)) seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val) loc_val = yield ctx.expr(1) loc = float(loc_val.item()) if self._is_tensor(loc_val) else float(loc_val) scale_val = yield ctx.expr(2) scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val) - shape_arg = self.shape - if len(ctx.expr()) > 3: - shape_arg = (yield ctx.expr(3)) generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Gumbel(loc, scale) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + return loc - scale * torch.log(-torch.log(torch.empty(shape_arg, device=self.device).uniform_(0, 1, generator=generator).clamp(min=1e-10)) + 1e-10) def visitWeibullDistFunc(self, ctx): seed_val = yield ctx.expr(0) + shape_arg = self.shape; + if len(ctx.expr()) > 3: + shape_arg = (yield ctx.expr(3)) seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val) scale_val = yield ctx.expr(1) scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val) @@ -1916,9 +1931,11 @@ class UnifiedMathVisitor(MathExprVisitor): shape_arg = self.shape if len(ctx.expr()) > 3: shape_arg = (yield ctx.expr(3)) + + # Implement Weibull using generator-aware uniform: scale * (-log(u))^(1/concentration) generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Weibull(scale, concentration) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + u = torch.rand(shape_arg, generator=generator, device=self.device) + return scale * torch.pow(-torch.log(u + 1e-10), 1.0 / concentration) def visitChi2DistFunc(self, ctx): seed_val = yield ctx.expr(0) @@ -1928,9 +1945,16 @@ class UnifiedMathVisitor(MathExprVisitor): shape_arg = self.shape if len(ctx.expr()) > 2: shape_arg = (yield ctx.expr(2)) - generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.Chi2(df) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + + # Chi-squared is Gamma(df/2, 2) + old_state = torch.get_rng_state() + try: + torch.manual_seed(seed) + dist = torch.distributions.Gamma(df / 2.0, 0.5) + result = dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])) + return result.to(device=self.device) + finally: + torch.set_rng_state(old_state) def visitStudentTDistFunc(self, ctx): seed_val = yield ctx.expr(0) @@ -1940,9 +1964,22 @@ class UnifiedMathVisitor(MathExprVisitor): shape_arg = self.shape if len(ctx.expr()) > 2: shape_arg = (yield ctx.expr(2)) + + # Student's t using normal and chi-squared: Z / sqrt(V/df) where Z~N(0,1) and V~Chi2(df) generator = torch.Generator(device=self.device).manual_seed(seed) - dist = torch.distributions.StudentT(df) - return dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])).to(device=self.device) + z = torch.randn(shape_arg, generator=generator, device=self.device) + + # Generate chi-squared using the same seed + 1 to maintain determinism but different samples + old_state = torch.get_rng_state() + try: + torch.manual_seed(seed + 1) + dist = torch.distributions.Gamma(df / 2.0, 0.5) + v = dist.sample(shape_arg if isinstance(shape_arg, torch.Size) else torch.Size(shape_arg) if isinstance(shape_arg, (list, tuple)) else torch.Size([shape_arg])) + v = v.to(device=self.device) + finally: + torch.set_rng_state(old_state) + + return z / torch.sqrt(v / df) def visitNvlFunc(self, ctx): v = yield ctx.expr(0) @@ -2098,12 +2135,12 @@ class UnifiedMathVisitor(MathExprVisitor): kh = kernel.view(1, kernel_size) - x_h = this._apply_conv_internal(x, kh, [kernel_size, 1], 2) + x_h = self._apply_conv_internal(x, kh, [kernel_size, 1], 2) kv = kernel.view(kernel_size, 1) - return this._apply_conv_internal(x_h, kv, [1, kernel_size], 2) + return self._apply_conv_internal(x_h, kv, [1, kernel_size], 2) - return this._apply_spatial_op(tsr, blur_op, original_shape) if reshap else blur_op(tsr) + return self._apply_spatial_op(tsr, blur_op, original_shape) if reshap else blur_op(tsr) def visitDistFunc(self, ctx): x1 = yield ctx.expr(0) @@ -2363,7 +2400,7 @@ class UnifiedMathVisitor(MathExprVisitor): # View tensors as integers if needed (bitwise ops require integer types) original_dtype_a = None original_dtype_b = None - + if self._is_tensor(a): original_dtype_a = a.dtype if a.dtype not in [torch.int8, torch.int16, torch.int32, torch.int64]: @@ -2371,7 +2408,7 @@ class UnifiedMathVisitor(MathExprVisitor): elem_size = a.element_size() view_dtype = self._get_bitwise_view_dtype(elem_size) a = a.view(view_dtype) - + if self._is_tensor(b): original_dtype_b = b.dtype if b.dtype not in [torch.int8, torch.int16, torch.int32, torch.int64]: @@ -2379,16 +2416,16 @@ class UnifiedMathVisitor(MathExprVisitor): elem_size = b.element_size() view_dtype = self._get_bitwise_view_dtype(elem_size) b = b.view(view_dtype) - + result = torch_op(a, b).contiguous() - + # View back to original dtype if we viewed a as non-integer if original_dtype_a is not None and original_dtype_a not in [torch.int8, torch.int16, torch.int32, torch.int64]: result = result.view(original_dtype_a) # View back to original dtype if we viewed b as non-integer (and didn't already view from a) elif original_dtype_b is not None and original_dtype_b not in [torch.int8, torch.int16, torch.int32, torch.int64]: result = result.view(original_dtype_b) - + return result.contiguous() return scalar_op(a, b) @@ -2401,7 +2438,7 @@ class UnifiedMathVisitor(MathExprVisitor): elem_size = t.element_size() if hasattr(t, 'element_size') else 4 view_dtype = self._get_bitwise_view_dtype(elem_size) original_dtype = t.dtype - + bits = t.view(view_dtype) res_bits = torch.bitwise_not(bits) return res_bits.view(original_dtype).contiguous() @@ -2448,10 +2485,10 @@ class UnifiedMathVisitor(MathExprVisitor): if counts.numel() == 1: return float(counts.item()) return counts - + if self._is_list(v): return [self._bitwise_popcount(x) for x in v] - + # Scalar - count set bits v_int = int(v) return float(bin(v_int & 0xFFFFFFFFFFFFFFFF).count('1')) @@ -2459,11 +2496,11 @@ class UnifiedMathVisitor(MathExprVisitor): def _scalar_bitwise_lshift(self, a, b): """Scalar left shift with bit-pattern preservation for floats.""" b_int = int(b) - + # If a is already an int, just do the shift if isinstance(a, int): return a << b_int - + # For floats, preserve bit pattern if isinstance(a, float): fmt = 'd' # double (64-bit) @@ -2474,18 +2511,18 @@ class UnifiedMathVisitor(MathExprVisitor): return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0] except struct.error: return float(result_bits & ((1 << 53) - 1)) # Return mantissa if error - + # Fallback for other types return int(a) << b_int def _scalar_bitwise_rshift(self, a, b): """Scalar right shift with bit-pattern preservation for floats.""" b_int = int(b) - + # If a is already an int, just do the shift if isinstance(a, int): return a >> b_int - + # For floats, preserve bit pattern if isinstance(a, float): fmt = 'd' # double (64-bit) @@ -2496,6 +2533,6 @@ class UnifiedMathVisitor(MathExprVisitor): return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0] except struct.error: return float(result_bits) - + # Fallback for other types return int(a) >> b_int diff --git a/more_math/Parser/optical_flow_utils.py b/more_math/Parser/optical_flow_utils.py index c298531..0d4b6e3 100644 --- a/more_math/Parser/optical_flow_utils.py +++ b/more_math/Parser/optical_flow_utils.py @@ -23,7 +23,7 @@ def warp(x, flow, padding_mode='reflection'): b, c, h, w = x.shape device = flow.device # Use flow device as the master device x = x.to(device) - + # Create grid grid_y, grid_x = torch.meshgrid( torch.linspace(0, h - 1, h, device=device), @@ -31,18 +31,18 @@ def warp(x, flow, padding_mode='reflection'): indexing='ij' ) # [H, W] - + grid = torch.stack((grid_x, grid_y), dim=0).unsqueeze(0).repeat(b, 1, 1, 1) # [B, 2, H, W] - + # Add flow to grid v_grid = grid + flow # [B, 2, H, W] - + # Normalize to [-1, 1] for grid_sample v_grid_x = 2.0 * v_grid[:, 0, :, :] / max(w - 1, 1) - 1.0 v_grid_y = 2.0 * v_grid[:, 1, :, :] / max(h - 1, 1) - 1.0 - + v_grid = torch.stack((v_grid_x, v_grid_y), dim=3) # [B, H, W, 2] with (x, y) - + return F.grid_sample(x, v_grid, mode='bilinear', padding_mode=padding_mode, align_corners=True) def preprocess_image(img, device): @@ -52,28 +52,28 @@ def preprocess_image(img, device): """ if img.ndim == 3: img = img.unsqueeze(0) - + # ComfyUI [B, H, W, C] -> Torch [B, C, H, W] img = img.movedim(-1, 1) - + # Ensure RGB if img.shape[1] == 1: img = img.expand(-1, 3, -1, -1) elif img.shape[1] > 3: img = img[:, :3, :, :] - + # Scale to [0, 255] as RAFT transforms usually expect this img = img * 255.0 - + # Use official transform if possible weights = Raft_Large_Weights.DEFAULT transform = weights.transforms() - + # The transform expects [0, 255] and returns normalized [-1, 1] # It takes (img1, img2) but we can use it for one or just follow its logic # Actually, let's just follow the logic: 2 * (img / 255.0) - 1.0 # Wait, if I do that, it's just 2 * img_01 - 1.0. - + return (2.0 * (img / 255.0) - 1.0).to(device) def get_optical_flow(img1, img2, tiling_size=0, iterations=12, multi_scale=False): @@ -83,23 +83,23 @@ def get_optical_flow(img1, img2, tiling_size=0, iterations=12, multi_scale=False """ device = img1.device model = get_raft_model(device) - + img1_proc = preprocess_image(img1, device) img2_proc = preprocess_image(img2, device) - + b, _, h, w = img1_proc.shape - + # Auto-tiling for large images to prevent OOM # 2MPx is a reasonable limit for 8GB VRAM without tiling if tiling_size <= 0 and (h * w) > 2000000: tiling_size = 1024 - + # Handle fraction or float conversion if 0 < tiling_size < 1: tiling_size = int(min(h, w) * tiling_size) else: tiling_size = int(tiling_size) - + # Ensure divisible by 8 for RAFT h_new, w_new = (h // 8) * 8, (w // 8) * 8 if h != h_new or w != w_new: @@ -112,7 +112,7 @@ def get_optical_flow(img1, img2, tiling_size=0, iterations=12, multi_scale=False for i in range(b): p1 = img1_proc[i:i+1] p2 = img2_proc[i:i+1] - + global_flow = None if multi_scale: # 1. Global Pass @@ -120,52 +120,52 @@ def get_optical_flow(img1, img2, tiling_size=0, iterations=12, multi_scale=False g_size = 256 g_h = (min(h_new, g_size) // 8) * 8 g_w = (min(w_new, g_size) // 8) * 8 - + gp1 = F.interpolate(p1, size=(g_h, g_w), mode='bilinear', align_corners=False) gp2 = F.interpolate(p2, size=(g_h, g_w), mode='bilinear', align_corners=False) - + # Global flow estimation global_flow = model(gp1, gp2, num_flow_updates=iterations)[-1] - + # Upsample global flow to high-res global_flow = F.interpolate(global_flow, size=(h_new, w_new), mode='bilinear', align_corners=False) global_flow[:, 0] *= float(w_new) / g_w global_flow[:, 1] *= float(h_new) / g_h - + # print(f"DEBUG: Global flow at center: {global_flow[0, :, h_new//2, w_new//2]}") - + # 2. Warp img2 using global flow p2_warped = warp(p2, global_flow) target_p2 = p2_warped else: target_p2 = p2 - + # 3. Local Refinement (Tiled or Full) if tiling_size > 0: f_residual = _tiled_flow(model, p1, target_p2, tiling_size, iterations) else: f_residual = model(p1, target_p2, num_flow_updates=iterations)[-1] - + # 4. Combine if global_flow is not None: total_flow = global_flow + f_residual else: total_flow = f_residual - + flows.append(total_flow) - + if i % 4 == 0 and b > 4: torch.cuda.empty_cache() - + flow = torch.cat(flows, dim=0) - + # Resize flow back if needed if h != h_new or w != w_new: flow = F.interpolate(flow, size=(h, w), mode='bilinear', align_corners=False) # Rescale flow values flow[:, 0] *= float(w) / w_new flow[:, 1] *= float(h) / h_new - + # [B, 2, H, W] -> [B, H, W, 2] return flow.movedim(1, -1) @@ -211,15 +211,15 @@ def apply_flow(image, flow): image = image.unsqueeze(0) if flow.ndim == 3: flow = flow.unsqueeze(0) - + # [B, H, W, 2] -> [B, 2, H, W] flow_nchw = flow.movedim(-1, 1) - + # Ensure image is on the same device as flow img_nchw = image.movedim(-1, 1).to(flow.device) - + warped = warp(img_nchw, flow_nchw) - + # [B, C, H, W] -> [B, H, W, C] return warped.movedim(1, -1) @@ -229,16 +229,16 @@ def flow_to_image(flow): """ if flow.ndim == 3: flow = flow.unsqueeze(0) - + # [B, H, W, 2] -> [B, 2, H, W] flow_torch = flow.movedim(-1, 1) - + # torchvision_flow_to_image expects (N, 2, H, W) img_uint8 = torchvision_flow_to_image(flow_torch) - + # [B, 3, H, W] -> [B, H, W, 3] img = img_uint8.movedim(1, -1).float() / 255.0 - + return img def _tiled_flow(model, img1, img2, tiling_size, iterations): @@ -247,53 +247,53 @@ def _tiled_flow(model, img1, img2, tiling_size, iterations): """ b, c, h, w = img1.shape device = img1.device - + # Enforce minimum tile size for RAFT if tiling_size < 256: # If user provided a very small value, maybe they meant fraction? # But for ComfyUI, we'll just enforce a safe minimum. tiling_size = max(256, tiling_size) - + # Ensure tiling_size is divisible by 8 tiling_size = (tiling_size // 8) * 8 - + # Clip tiling_size to image dimensions t_h = int(min(tiling_size, h)) t_w = int(min(tiling_size, w)) - + # If image is smaller than tiling_size after divisibility adjustments if t_h < 8 or t_w < 8: return model(img1, img2, num_flow_updates=iterations)[-1] output_flow = torch.zeros((b, 2, h, w), device=device) weights = torch.zeros((b, 1, h, w), device=device) - + # Stride of 50% overlap is usually good stride_h = int(t_h // 2) stride_w = int(t_w // 2) - + # Generate coordinates for tiles y_coords = list(range(0, h - t_h, stride_h)) + [max(0, h - t_h)] x_coords = list(range(0, w - t_w, stride_w)) + [max(0, w - t_w)] - + # Remove duplicates if any (when image size matches stride) y_coords = sorted(list(set(y_coords))) x_coords = sorted(list(set(x_coords))) - + for y in y_coords: for x in x_coords: patch1 = img1[:, :, y:y+t_h, x:x+t_w] patch2 = img2[:, :, y:y+t_h, x:x+t_w] - + # RAFT expects inputs divisible by 8, but we already ensured t_h, t_w are. patch_flow = model(patch1, patch2, num_flow_updates=iterations)[-1] - + # Create local mask for this tile # Linear tapering only on overlapping edges layer_mask = torch.ones((1, 1, t_h, t_w), device=device) overlap_h = stride_h overlap_w = stride_w - + if y > 0: # Fade in top layer_mask[:, :, :overlap_h, :] *= torch.linspace(0, 1, overlap_h, device=device).view(1, 1, overlap_h, 1) if y + t_h < h: # Fade out bottom @@ -302,8 +302,8 @@ def _tiled_flow(model, img1, img2, tiling_size, iterations): layer_mask[:, :, :, :overlap_w] *= torch.linspace(0, 1, overlap_w, device=device).view(1, 1, 1, overlap_w) if x + t_w < w: # Fade out right layer_mask[:, :, :, -overlap_w:] *= torch.linspace(1, 0, overlap_w, device=device).view(1, 1, 1, overlap_w) - + output_flow[:, :, y:y+t_h, x:x+t_w] += patch_flow * layer_mask weights[:, :, y:y+t_h, x:x+t_w] += layer_mask - + return output_flow / torch.clamp(weights, min=1e-6)