Better struffies

More resolutions and radial attention calcs
This commit is contained in:
ImpactFrames
2025-08-25 15:07:53 +01:00
committed by GitHub
parent a2fe6fa45c
commit d7b8369f32
+238 -112
View File
@@ -1,3 +1,145 @@
import math
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
RADIAL_ALIGNMENT = VAE_STRIDE[1] * PATCH_SIZE[1] # 16 for height/width
PREFERED_KONTEXT_RESOLUTIONS = [
(672, 1568),
(688, 1504),
(720, 1456),
(752, 1392),
(800, 1328),
(832, 1248),
(880, 1184),
(944, 1104),
(1024, 1024),
(1104, 944),
(1184, 880),
(1248, 832),
(1328, 800),
(1392, 752),
(1456, 720),
(1504, 688),
(1568, 672),
]
def calculate_radial_compatible_resolution(width, height, mode="closest", block_size=64):
"""
Calculate radial attention compatible resolution.
For radial attention to work, patches_per_frame must be divisible by block_size:
patches_per_frame = (height//8) * (width//8) // 4
Args:
width (int): Original width
height (int): Original height
mode (str): "upscale", "downscale", or "closest"
block_size (int): Radial attention block size (64 or 128)
Returns:
tuple: (compatible_width, compatible_height)
"""
def find_compatible_dimension(target_size, mode, block_size):
# Start with VAE-aligned size
base_size = (target_size // VAE_STRIDE[1]) * VAE_STRIDE[1]
# Search for a size where patches_per_frame % block_size == 0
search_range = range(max(VAE_STRIDE[1], base_size - 64), base_size + 80, VAE_STRIDE[1])
candidates = []
for test_size in search_range:
lat_dim = test_size // VAE_STRIDE[1]
patches_per_frame = (lat_dim * lat_dim) // (PATCH_SIZE[1] * PATCH_SIZE[2])
if patches_per_frame % block_size == 0:
distance = abs(test_size - target_size)
candidates.append((distance, test_size))
if not candidates:
# Fallback: just ensure VAE alignment
return base_size if base_size >= VAE_STRIDE[1] else VAE_STRIDE[1]
candidates.sort() # Sort by distance
if mode == "upscale":
valid_candidates = [size for dist, size in candidates if size >= target_size]
return valid_candidates[0] if valid_candidates else candidates[-1][1]
elif mode == "downscale":
valid_candidates = [size for dist, size in candidates if size <= target_size]
return valid_candidates[0] if valid_candidates else candidates[0][1]
else: # closest
return candidates[0][1]
# Handle square resolutions specially (both dimensions must work together)
if width == height:
# For square resolutions, find size where lat_dim^2 // 4 % block_size == 0
target_lat = width // VAE_STRIDE[1]
for offset in range(-8, 9): # Search around target
test_lat = target_lat + offset
if test_lat <= 0:
continue
patches_per_frame = (test_lat * test_lat) // (PATCH_SIZE[1] * PATCH_SIZE[2])
if patches_per_frame % block_size == 0:
test_size = test_lat * VAE_STRIDE[1]
distance = abs(test_size - width)
if mode == "upscale" and test_size >= width:
return test_size, test_size
elif mode == "downscale" and test_size <= width:
return test_size, test_size
elif mode == "closest":
# Check if this is closer than the original
orig_lat = width // VAE_STRIDE[1]
orig_patches = (orig_lat * orig_lat) // 4
if orig_patches % block_size != 0 or distance == 0:
return test_size, test_size
# If no perfect match found, use the original if it's already compatible
orig_lat = width // VAE_STRIDE[1]
orig_patches = (orig_lat * orig_lat) // 4
if orig_patches % block_size == 0:
return width, height
# For non-square or fallback, handle dimensions independently
compatible_width = find_compatible_dimension(width, mode, block_size)
compatible_height = find_compatible_dimension(height, mode, block_size)
return compatible_width, compatible_height
def update_resolutions_for_radial_attention(resolutions_dict, mode="closest", block_size=64):
"""
Update resolution dictionary to make all resolutions radial attention compatible.
Args:
resolutions_dict (dict): Original resolutions dictionary
mode (str): "upscale", "downscale", or "closest"
block_size (int): Radial attention block size (64 or 128)
Returns:
dict: Updated resolutions dictionary
"""
updated_resolutions = {}
for model_type, orientations in resolutions_dict.items():
updated_resolutions[model_type] = {}
for orientation, qualities in orientations.items():
updated_resolutions[model_type][orientation] = {}
for quality, (width, height) in qualities.items():
new_width, new_height = calculate_radial_compatible_resolution(width, height, mode, block_size)
updated_resolutions[model_type][orientation][quality] = (new_width, new_height)
# Log changes if resolution was modified
if new_width != width or new_height != height:
print(f"Radial Attention (block_size={block_size}): {model_type}-{orientation}-{quality}: {width}x{height} -> {new_width}x{new_height}")
return updated_resolutions
class VideoResolutionSelector:
"""
Selects appropriate video resolution based on mode, aspect ratio, and quality settings.
@@ -7,91 +149,60 @@ class VideoResolutionSelector:
# Resolution mappings based on mode, aspect ratio, and quality
RESOLUTIONS = {
"I2V720p": {
"Horizontal": {
"HQ": (1280, 720),
"MQ": (832, 480),
"LQ": (704, 544)
},
"Vertical": {
"HQ": (720, 1280),
"MQ": (480, 832),
"LQ": (544, 704)
},
"Squarish": {
"HQ": (624, 624),
"MQ": (624, 624),
"LQ": (624, 624)
}
"Horizontal": {"HQ": (1280, 720), "MQ": (832, 480), "LQ": (704, 544)},
"Vertical": {"HQ": (720, 1280), "MQ": (480, 832), "LQ": (544, 704)},
"Squarish": {"HQ": (624, 624), "MQ": (624, 624), "LQ": (624, 624)},
},
"I2V480p": {
"Horizontal": {
"HQ": (832, 480),
"MQ": (704, 544),
"LQ": (704, 544)
},
"Vertical": {
"HQ": (480, 832),
"MQ": (544, 704),
"LQ": (544, 704)
},
"Squarish": {
"HQ": (624, 624),
"MQ": (624, 624),
"LQ": (624, 624)
}
"Horizontal": {"HQ": (832, 480), "MQ": (704, 544), "LQ": (704, 544)},
"Vertical": {"HQ": (480, 832), "MQ": (544, 704), "LQ": (544, 704)},
"Squarish": {"HQ": (624, 624), "MQ": (624, 624), "LQ": (624, 624)},
},
"T2V14B": {
"Horizontal": {
"HQ": (1280, 720),
"MQ": (1088, 832),
"LQ": (832, 480)
},
"Vertical": {
"HQ": (720, 1280),
"MQ": (832, 1088),
"LQ": (480, 832)
},
"Squarish": {
"HQ": (960, 960),
"MQ": (624, 624),
"LQ": (544, 704)
}
"Horizontal": {"HQ": (1280, 720), "MQ": (1088, 832), "LQ": (832, 480)},
"Vertical": {"HQ": (720, 1280), "MQ": (832, 1088), "LQ": (480, 832)},
"Squarish": {"HQ": (960, 960), "MQ": (624, 624), "LQ": (544, 704)},
},
"T2V1.3B": {
"Horizontal": {
"HQ": (832, 480),
"MQ": (704, 544),
"LQ": (704, 544)
},
"Vertical": {
"HQ": (480, 832),
"MQ": (544, 704),
"LQ": (544, 704)
},
"Squarish": {
"HQ": (624, 624),
"MQ": (624, 624),
"LQ": (624, 624)
}
}
"Horizontal": {"HQ": (832, 480), "MQ": (704, 544), "LQ": (704, 544)},
"Vertical": {"HQ": (480, 832), "MQ": (544, 704), "LQ": (544, 704)},
"Squarish": {"HQ": (624, 624), "MQ": (624, 624), "LQ": (624, 624)},
},
"IMG": {
"Horizontal": {"HQ": (1600, 900), "MQ": (1280, 720), "LQ": (1024, 576)},
"Vertical": {"HQ": (900, 1600), "MQ": (720, 1280), "LQ": (576, 1024)},
"Squarish": {"HQ": (1600, 1600), "MQ": (1024, 1024), "LQ": (512, 512)},
"Cinematic": {"HQ": (1600, 688), "MQ": (1280, 550), "LQ": (1024, 440)}, # ≈2.35:1
},
"KONTEXT": {
"Vertical": {"HQ": (672, 1568), "MQ": (720, 1456), "LQ": (832, 1248)},
"Horizontal": {"HQ": (1568, 672), "MQ": (1456, 720), "LQ": (1248, 832)},
"Squarish": {"HQ": (1024, 1024), "MQ": (944, 1104), "LQ": (880, 1184)},
},
"QWEN": {
"Square": {"HQ": (1024, 1024), "MQ": (768, 768), "LQ": (512, 512)},
"Landscape": {"HQ": (1280, 720), "MQ": (1024, 768), "LQ": (832, 624)},
"Portrait": {"HQ": (720, 1280), "MQ": (768, 1024), "LQ": (624, 832)},
"Wide": {"HQ": (1536, 768), "MQ": (1280, 640), "LQ": (1024, 512)},
"Tall": {"HQ": (768, 1536), "MQ": (640, 1280), "LQ": (512, 1024)},
"UltraWide": {"HQ": (1792, 768), "MQ": (1536, 640), "LQ": (1280, 544)},
"UltraTall": {"HQ": (768, 1792), "MQ": (640, 1536), "LQ": (544, 1280)},
},
}
@classmethod
def INPUT_TYPES(cls):
modes = list(cls.RESOLUTIONS.keys())
return {
"required": {
"mode": (["I2V480p", "I2V720p", "T2V1.3B", "T2V14B"], {
"default": "I2V720p",
"tooltip": "Select the video generation mode"
}),
"aspect_ratio": (["Horizontal", "Vertical", "Squarish"], {
"default": "Horizontal",
"tooltip": "Select the aspect ratio orientation"
}),
"quality": (["HQ", "MQ", "LQ"], {
"default": "HQ",
"tooltip": "Select quality level - HQ: High Quality, MQ: Medium Quality, LQ: Low Quality"
}),
"mode": (modes, {"default": modes[0], "tooltip": "Generation mode"}),
"aspect_ratio": (["Horizontal", "Vertical", "Squarish", "Cinematic", "Square", "Landscape", "Portrait", "Wide", "Tall", "UltraWide", "UltraTall"], {"default": "Horizontal"}),
"quality": (["HQ", "MQ", "LQ"], {"default": "HQ"}),
},
"optional": {
"enable_radial_attention": ("BOOLEAN", {"default": False, "tooltip": "Enable radial attention compatibility"}),
"radial_mode": (["upscale", "downscale", "closest"], {"default": "upscale", "tooltip": "How to adjust resolutions for radial attention"}),
"block_size": ([64, 128], {"default": 64, "tooltip": "Radial attention block size"}),
}
}
@@ -100,46 +211,17 @@ class VideoResolutionSelector:
FUNCTION = "get_resolution"
CATEGORY = "ImpactFrames💥🎞️/utils"
DESCRIPTION = """
Automatically selects the appropriate width and height based on video generation mode,
aspect ratio, and quality settings. Compatible with KJNodes image resize nodes.
Modes:
- I2V480p: Image to Video 480p
- I2V720p: Image to Video 720p
- T2V1.3B: Text to Video 1.3B model
- T2V14B: Text to Video 14B model
Aspect Ratios:
- Horizontal: Wider than tall (landscape)
- Vertical: Taller than wide (portrait)
- Squarish: Roughly square aspect ratio
Quality:
- HQ: Highest available resolution for the mode
- MQ: Medium quality/resolution
- LQ: Lower quality/resolution
"""
def get_resolution(self, mode, aspect_ratio, quality):
"""
Returns the width and height based on the selected parameters.
Args:
mode: The video generation mode
aspect_ratio: The desired aspect ratio
quality: The quality level
Returns:
tuple: (width, height)
"""
def get_resolution(self, mode, aspect_ratio, quality, enable_radial_attention=False, radial_mode="upscale", block_size=64):
try:
width, height = self.RESOLUTIONS[mode][aspect_ratio][quality]
return (width, height)
w, h = self.RESOLUTIONS[mode][aspect_ratio][quality]
# Apply radial attention compatibility if enabled
if enable_radial_attention:
w, h = calculate_radial_compatible_resolution(w, h, radial_mode, block_size)
return (w, h)
except KeyError:
# Fallback to a default resolution if something goes wrong
print(f"Warning: Invalid combination of mode={mode}, aspect_ratio={aspect_ratio}, quality={quality}")
print("Falling back to default resolution 832x480")
# fallback default
return (832, 480)
@@ -150,4 +232,48 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"VideoResolutionSelector": "Video Resolution Selector 🎬",
}
}
def get_radial_resolutions(mode="closest", block_size=64):
"""Get radial attention compatible resolutions."""
return update_resolutions_for_radial_attention(VideoResolutionSelector.RESOLUTIONS, mode, block_size)
if __name__ == "__main__":
# Test the calculator
print("Testing radial attention compatible resolutions:")
for block_size in [64, 128]:
print(f"\n=== BLOCK SIZE {block_size} ===")
print(f"Original problematic size: 624x624")
print("Upscale mode:", calculate_radial_compatible_resolution(624, 624, "upscale", block_size))
print("Downscale mode:", calculate_radial_compatible_resolution(624, 624, "downscale", block_size))
print("Closest mode:", calculate_radial_compatible_resolution(624, 624, "closest", block_size))
print(f"\n--- Resolution Sets (block_size={block_size}) ---")
for mode in ["upscale", "downscale", "closest"]:
print(f"\n{mode.upper()} MODE:")
radial_resolutions = get_radial_resolutions(mode, block_size)
# Show squarish resolutions (most affected)
print("Squarish resolutions:")
for model_type in radial_resolutions:
if "Squarish" in radial_resolutions[model_type]:
for quality, (w, h) in radial_resolutions[model_type]["Squarish"].items():
original = VideoResolutionSelector.RESOLUTIONS[model_type]["Squarish"][quality]
changed = "✓" if (w, h) != original else " "
print(f" {model_type}-{quality}: {w}x{h} {changed}")
# Test a few other problematic ones
print("Other potentially problematic:")
test_cases = [
("T2V14B", "Horizontal", "MQ"), # 1088x832
("IMG", "Cinematic", "LQ"), # 1024x440
]
for model, orient, qual in test_cases:
if orient in radial_resolutions[model] and qual in radial_resolutions[model][orient]:
w, h = radial_resolutions[model][orient][qual]
original = VideoResolutionSelector.RESOLUTIONS[model][orient][qual]
changed = "✓" if (w, h) != original else " "
print(f" {model}-{orient}-{qual}: {w}x{h} {changed}")