ruff + AI: fix determnism

This commit is contained in:
mcDandy
2026-02-14 23:04:17 +01:00
parent f57f224fd6
commit 7cd8e0362e
4 changed files with 168 additions and 132 deletions
-1
View File
@@ -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
+38 -38
View File
@@ -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)
+79 -42
View File
@@ -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
+51 -51
View File
@@ -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)