mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-05 07:07:05 +08:00
vae tiled fixes and few other mistakes
This commit is contained in:
parent
b9f5145f4d
commit
aaed282c3a
@ -800,7 +800,7 @@ class LoadedModel:
|
|||||||
return self.model_offloaded_memory()
|
return self.model_offloaded_memory()
|
||||||
|
|
||||||
# Handle AutoencoderKL
|
# Handle AutoencoderKL
|
||||||
if self.model.model is not None and isinstance(self.model.model, AutoencoderKL):
|
if self.model is not None and isinstance(self.model.model, AutoencoderKL):
|
||||||
shape = getattr(self.model, 'last_shape', (1, 4, 64, 64))
|
shape = getattr(self.model, 'last_shape', (1, 4, 64, 64))
|
||||||
dtype = getattr(self.model, 'model_dtype', torch.float32)()
|
dtype = getattr(self.model, 'model_dtype', torch.float32)()
|
||||||
return estimate_vae_decode_memory(self.model.model, shape, dtype)
|
return estimate_vae_decode_memory(self.model.model, shape, dtype)
|
||||||
|
|||||||
40
comfy/sd.py
40
comfy/sd.py
@ -28,6 +28,7 @@ from . import model_detection
|
|||||||
|
|
||||||
from . import sd1_clip
|
from . import sd1_clip
|
||||||
from . import sdxl_clip
|
from . import sdxl_clip
|
||||||
|
from comfy.cli_args import args
|
||||||
import comfy.text_encoders.sd2_clip
|
import comfy.text_encoders.sd2_clip
|
||||||
import comfy.text_encoders.sd3_clip
|
import comfy.text_encoders.sd3_clip
|
||||||
import comfy.text_encoders.sa_t5
|
import comfy.text_encoders.sa_t5
|
||||||
@ -54,6 +55,8 @@ import comfy.taesd.taesd
|
|||||||
|
|
||||||
import comfy.ldm.flux.redux
|
import comfy.ldm.flux.redux
|
||||||
|
|
||||||
|
DEBUG_ENABLED = args.debug
|
||||||
|
|
||||||
def load_lora_for_models(model, clip, lora, strength_model, strength_clip):
|
def load_lora_for_models(model, clip, lora, strength_model, strength_clip):
|
||||||
key_map = {}
|
key_map = {}
|
||||||
if model is not None:
|
if model is not None:
|
||||||
@ -497,18 +500,23 @@ class VAE:
|
|||||||
pixels = pixels.narrow(d + 1, x_offset, x)
|
pixels = pixels.narrow(d + 1, x_offset, x)
|
||||||
return pixels
|
return pixels
|
||||||
|
|
||||||
def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16):
|
def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap=16):
|
||||||
|
# Calculate progress bar steps for a single pass
|
||||||
steps = samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x, tile_y, overlap)
|
steps = samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x, tile_y, overlap)
|
||||||
steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x // 2, tile_y * 2, overlap)
|
|
||||||
steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x * 2, tile_y // 2, overlap)
|
|
||||||
pbar = comfy.utils.ProgressBar(steps)
|
pbar = comfy.utils.ProgressBar(steps)
|
||||||
|
|
||||||
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float()
|
# Define decode function with tile logging
|
||||||
output = self.process_output(
|
if not DEBUG_ENABLED:
|
||||||
(comfy.utils.tiled_scale(samples, decode_fn, tile_x // 2, tile_y * 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) +
|
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float()
|
||||||
comfy.utils.tiled_scale(samples, decode_fn, tile_x * 2, tile_y // 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) +
|
else:
|
||||||
comfy.utils.tiled_scale(samples, decode_fn, tile_x, tile_y, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar))
|
decode_fn = lambda a: (logging.debug(f"Tile shape: {a.shape}, min: {a.min()}, max: {a.max()}"),
|
||||||
/ 3.0)
|
self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float())[1]
|
||||||
|
|
||||||
|
# Single pass with provided tile sizes
|
||||||
|
output = comfy.utils.tiled_scale(
|
||||||
|
samples, decode_fn, tile_x, tile_y, overlap,
|
||||||
|
upscale_amount=self.upscale_ratio, output_device=self.output_device, pbar=pbar
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
def decode_tiled_1d(self, samples, tile_x=128, overlap=32):
|
def decode_tiled_1d(self, samples, tile_x=128, overlap=32):
|
||||||
@ -1068,9 +1076,10 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
|
|||||||
else:
|
else:
|
||||||
logging.warning("no CLIP/text encoder weights in checkpoint, the text encoder model will not be loaded.")
|
logging.warning("no CLIP/text encoder weights in checkpoint, the text encoder model will not be loaded.")
|
||||||
|
|
||||||
left_over = sd.keys()
|
if DEBUG_ENABLED:
|
||||||
if len(left_over) > 0:
|
left_over = sd.keys()
|
||||||
logging.debug("left over keys: {}".format(left_over))
|
if len(left_over) > 0:
|
||||||
|
logging.debug("left over keys: {}".format(left_over))
|
||||||
|
|
||||||
if output_model:
|
if output_model:
|
||||||
model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=model_management.unet_offload_device())
|
model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=model_management.unet_offload_device())
|
||||||
@ -1137,9 +1146,10 @@ def load_diffusion_model_state_dict(sd, model_options={}): #load unet in diffuse
|
|||||||
model = model_config.get_model(new_sd, "")
|
model = model_config.get_model(new_sd, "")
|
||||||
model = model.to(offload_device)
|
model = model.to(offload_device)
|
||||||
model.load_model_weights(new_sd, "")
|
model.load_model_weights(new_sd, "")
|
||||||
left_over = sd.keys()
|
if DEBUG_ENABLED:
|
||||||
if len(left_over) > 0:
|
left_over = sd.keys()
|
||||||
logging.info("left over keys in unet: {}".format(left_over))
|
if len(left_over) > 0:
|
||||||
|
logging.info("left over keys in unet: {}".format(left_over))
|
||||||
return comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device)
|
return comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
198
comfy/utils.py
198
comfy/utils.py
@ -33,6 +33,9 @@ from comfy.cli_args import args
|
|||||||
MMAP_TORCH_FILES = args.mmap_torch_files
|
MMAP_TORCH_FILES = args.mmap_torch_files
|
||||||
|
|
||||||
ALWAYS_SAFE_LOAD = False
|
ALWAYS_SAFE_LOAD = False
|
||||||
|
|
||||||
|
DEBUG_ENABLED = args.debug
|
||||||
|
|
||||||
if hasattr(torch.serialization, "add_safe_globals"): # TODO: this was added in pytorch 2.4, the unsafe path should be removed once earlier versions are deprecated
|
if hasattr(torch.serialization, "add_safe_globals"): # TODO: this was added in pytorch 2.4, the unsafe path should be removed once earlier versions are deprecated
|
||||||
class ModelCheckpoint:
|
class ModelCheckpoint:
|
||||||
pass
|
pass
|
||||||
@ -876,116 +879,205 @@ def get_tiled_scale_steps(width, height, tile_x, tile_y, overlap):
|
|||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", downscale=False, index_formulas=None, pbar=None):
|
def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", downscale=False, index_formulas=None, pbar=None):
|
||||||
|
"""
|
||||||
|
Perform tiled scaling of input samples using the provided function with overlap blending.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
samples: Input tensor of shape [batch, channels, *spatial_dims].
|
||||||
|
function: Function to process each tile (e.g., VAE decode).
|
||||||
|
tile: Tuple of tile sizes for each spatial dimension.
|
||||||
|
overlap: Overlap size or list of overlaps for each dimension.
|
||||||
|
upscale_amount: Scaling factor or list of factors for each dimension.
|
||||||
|
out_channels: Number of output channels.
|
||||||
|
output_device: Device for output tensor.
|
||||||
|
downscale: If True, downscale instead of upscale.
|
||||||
|
index_formulas: Scaling factors for tile positions (defaults to upscale_amount).
|
||||||
|
pbar: Optional progress bar object.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Scaled output tensor of shape [batch, out_channels, *scaled_spatial_dims].
|
||||||
|
"""
|
||||||
dims = len(tile)
|
dims = len(tile)
|
||||||
|
|
||||||
if not (isinstance(upscale_amount, (tuple, list))):
|
# Wrap function to ensure FP32
|
||||||
|
def fp32_function(x):
|
||||||
|
x_fp32 = x.to(dtype=torch.float32)
|
||||||
|
result = function(x_fp32)
|
||||||
|
return result.to(dtype=torch.float32)
|
||||||
|
|
||||||
|
# Convert parameters to lists for multidimensional support
|
||||||
|
if not isinstance(upscale_amount, (tuple, list)):
|
||||||
upscale_amount = [upscale_amount] * dims
|
upscale_amount = [upscale_amount] * dims
|
||||||
|
if not isinstance(overlap, (tuple, list)):
|
||||||
if not (isinstance(overlap, (tuple, list))):
|
|
||||||
overlap = [overlap] * dims
|
overlap = [overlap] * dims
|
||||||
|
|
||||||
if index_formulas is None:
|
if index_formulas is None:
|
||||||
index_formulas = upscale_amount
|
index_formulas = upscale_amount
|
||||||
|
if not isinstance(index_formulas, (tuple, list)):
|
||||||
if not (isinstance(index_formulas, (tuple, list))):
|
|
||||||
index_formulas = [index_formulas] * dims
|
index_formulas = [index_formulas] * dims
|
||||||
|
|
||||||
|
# Define scaling functions
|
||||||
def get_upscale(dim, val):
|
def get_upscale(dim, val):
|
||||||
up = upscale_amount[dim]
|
up = upscale_amount[dim]
|
||||||
if callable(up):
|
return up(val) if callable(up) else up * val
|
||||||
return up(val)
|
|
||||||
else:
|
|
||||||
return up * val
|
|
||||||
|
|
||||||
def get_downscale(dim, val):
|
def get_downscale(dim, val):
|
||||||
up = upscale_amount[dim]
|
up = upscale_amount[dim]
|
||||||
if callable(up):
|
return up(val) if callable(up) else val / up
|
||||||
return up(val)
|
|
||||||
else:
|
|
||||||
return val / up
|
|
||||||
|
|
||||||
def get_upscale_pos(dim, val):
|
def get_upscale_pos(dim, val):
|
||||||
up = index_formulas[dim]
|
up = index_formulas[dim]
|
||||||
if callable(up):
|
return up(val) if callable(up) else up * val
|
||||||
return up(val)
|
|
||||||
else:
|
|
||||||
return up * val
|
|
||||||
|
|
||||||
def get_downscale_pos(dim, val):
|
def get_downscale_pos(dim, val):
|
||||||
up = index_formulas[dim]
|
up = index_formulas[dim]
|
||||||
if callable(up):
|
return up(val) if callable(up) else val / up
|
||||||
return up(val)
|
|
||||||
else:
|
|
||||||
return val / up
|
|
||||||
|
|
||||||
if downscale:
|
get_scale = get_downscale if downscale else get_upscale
|
||||||
get_scale = get_downscale
|
get_pos = get_downscale_pos if downscale else get_upscale_pos
|
||||||
get_pos = get_downscale_pos
|
|
||||||
else:
|
|
||||||
get_scale = get_upscale
|
|
||||||
get_pos = get_upscale_pos
|
|
||||||
|
|
||||||
def mult_list_upscale(a):
|
def mult_list_upscale(a):
|
||||||
out = []
|
"""Compute scaled dimensions for output tensor."""
|
||||||
for i in range(len(a)):
|
return [round(get_scale(i, a[i])) for i in range(len(a))]
|
||||||
out.append(round(get_scale(i, a[i])))
|
|
||||||
return out
|
|
||||||
|
|
||||||
output = torch.empty([samples.shape[0], out_channels] + mult_list_upscale(samples.shape[2:]), device=output_device)
|
# Initialize output tensor
|
||||||
|
output_shape = [samples.shape[0], out_channels] + mult_list_upscale(samples.shape[2:])
|
||||||
|
output = torch.empty(output_shape, device=output_device, dtype=torch.float32)
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"Input shape: {samples.shape}, output shape: {output_shape}, tile: {tile}")
|
||||||
|
logging.debug(f"Input stats: min={samples.min():.4f}, max={samples.max():.4f}, mean={samples.mean():.4f}")
|
||||||
|
|
||||||
|
# Test VAE
|
||||||
|
try:
|
||||||
|
test_input = samples[:1].to(dtype=torch.float32)
|
||||||
|
test_result = fp32_function(test_input).to(output_device, dtype=torch.float32)
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"VAE test result: shape={test_result.shape}, min={test_result.min():.4f}, max={test_result.max():.4f}")
|
||||||
|
if torch.isnan(test_result).any() or torch.isinf(test_result).any():
|
||||||
|
logging.error("VAE produces NaN or Inf in test output. Check VAE model or input latents.")
|
||||||
|
raise RuntimeError("VAE output contains NaN or Inf")
|
||||||
|
except Exception as e:
|
||||||
|
logging.error(f"VAE test failed: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
for b in range(samples.shape[0]):
|
for b in range(samples.shape[0]):
|
||||||
s = samples[b:b+1]
|
s = samples[b:b+1]
|
||||||
|
|
||||||
# handle entire input fitting in a single tile
|
# Handle case where input fits in a single tile
|
||||||
if all(s.shape[d+2] <= tile[d] for d in range(dims)):
|
if all(s.shape[d+2] <= tile[d] for d in range(dims)):
|
||||||
output[b:b+1] = function(s).to(output_device)
|
s_fp32 = s.to(dtype=torch.float32)
|
||||||
|
result = fp32_function(s_fp32).to(output_device, dtype=torch.float32)
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"Single tile result: shape={result.shape}, min={result.min():.4f}, max={result.max():.4f}")
|
||||||
|
if result.shape == output_shape[1:]:
|
||||||
|
output[b:b+1] = result
|
||||||
|
else:
|
||||||
|
result = result.narrow(1, 0, output_shape[1])
|
||||||
|
for d in range(dims):
|
||||||
|
result = result.narrow(d + 2, 0, output_shape[d + 2])
|
||||||
|
output[b:b+1] = result
|
||||||
if pbar is not None:
|
if pbar is not None:
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
out = torch.zeros([s.shape[0], out_channels] + mult_list_upscale(s.shape[2:]), device=output_device)
|
# Initialize accumulation tensors
|
||||||
out_div = torch.zeros([s.shape[0], out_channels] + mult_list_upscale(s.shape[2:]), device=output_device)
|
out = torch.zeros(output_shape[1:], device=output_device, dtype=torch.float32)
|
||||||
|
out_div = torch.full_like(out, 1e-6)
|
||||||
|
|
||||||
positions = [range(0, s.shape[d+2] - overlap[d], tile[d] - overlap[d]) if s.shape[d+2] > tile[d] else [0] for d in range(dims)]
|
# Compute tile positions
|
||||||
|
positions = []
|
||||||
|
tile_counts = []
|
||||||
|
for d in range(dims):
|
||||||
|
step = max(1, tile[d] - overlap[d])
|
||||||
|
end = max(0, s.shape[d+2] - tile[d])
|
||||||
|
pos = list(range(0, end + 1, step))
|
||||||
|
if pos and (pos[-1] < end or s.shape[d+2] > tile[d]):
|
||||||
|
pos.append(end)
|
||||||
|
positions.append(pos if pos else [0])
|
||||||
|
tile_counts.append(len(pos))
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"Tile positions: {positions}")
|
||||||
|
|
||||||
|
# Process each tile
|
||||||
|
total_tiles = max(1, len(list(itertools.product(*positions))))
|
||||||
for it in itertools.product(*positions):
|
for it in itertools.product(*positions):
|
||||||
s_in = s
|
s_in = s
|
||||||
upscaled = []
|
upscaled = []
|
||||||
|
|
||||||
|
# Extract tile
|
||||||
for d in range(dims):
|
for d in range(dims):
|
||||||
pos = max(0, min(s.shape[d + 2] - overlap[d], it[d]))
|
pos = max(0, min(s.shape[d+2] - tile[d], it[d]))
|
||||||
l = min(tile[d], s.shape[d + 2] - pos)
|
length = min(tile[d], s.shape[d+2] - pos)
|
||||||
s_in = s_in.narrow(d + 2, pos, l)
|
s_in = s_in.narrow(d + 2, pos, length)
|
||||||
upscaled.append(round(get_pos(d, pos)))
|
upscaled.append(round(get_pos(d, pos)))
|
||||||
|
|
||||||
ps = function(s_in).to(output_device)
|
# Process tile
|
||||||
mask = torch.ones_like(ps)
|
s_in_fp32 = s_in.to(dtype=torch.float32)
|
||||||
|
ps = fp32_function(s_in_fp32).to(output_device, dtype=torch.float32)
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"Tile at {it}: input={s_in.shape}, output={ps.shape}, min={ps.min():.4f}, max={ps.max():.4f}")
|
||||||
|
if torch.isnan(ps).any() or torch.isinf(ps).any():
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.warning(f"Tile at {it} contains NaN or Inf, clamping values")
|
||||||
|
ps = torch.clamp(ps, min=-1e6, max=1e6)
|
||||||
|
mask = torch.ones_like(ps, dtype=torch.float32) / total_tiles
|
||||||
|
|
||||||
for d in range(2, dims + 2):
|
# Apply feathering for smooth overlap blending
|
||||||
feather = round(get_scale(d - 2, overlap[d - 2]))
|
for d in range(dims):
|
||||||
if feather >= mask.shape[d]:
|
feather = min(round(get_scale(d, overlap[d])), ps.shape[d + 2] // 16)
|
||||||
|
if feather < 1:
|
||||||
continue
|
continue
|
||||||
for t in range(feather):
|
for t in range(feather):
|
||||||
a = (t + 1) / feather
|
a = (t + 1) / feather
|
||||||
mask.narrow(d, t, 1).mul_(a)
|
mask.narrow(d + 2, t, 1).mul_(a)
|
||||||
mask.narrow(d, mask.shape[d] - 1 - t, 1).mul_(a)
|
mask.narrow(d + 2, mask.shape[d + 2] - 1 - t, 1).mul_(a)
|
||||||
|
mask = mask.clamp(min=1e-6)
|
||||||
|
|
||||||
|
# Accumulate results
|
||||||
o = out
|
o = out
|
||||||
o_d = out_div
|
o_d = out_div
|
||||||
for d in range(dims):
|
for d in range(dims):
|
||||||
o = o.narrow(d + 2, upscaled[d], mask.shape[d + 2])
|
start = upscaled[d]
|
||||||
o_d = o_d.narrow(d + 2, upscaled[d], mask.shape[d + 2])
|
size = min(ps.shape[d + 2], output_shape[d + 2] - start)
|
||||||
|
o = o.narrow(d + 1, start, size)
|
||||||
|
o_d = o_d.narrow(d + 1, start, size)
|
||||||
|
ps = ps.narrow(d + 2, 0, size)
|
||||||
|
mask = mask.narrow(d + 2, 0, size)
|
||||||
|
|
||||||
|
# Squeeze batch dimension
|
||||||
|
ps = ps.squeeze(0)
|
||||||
|
mask = mask.squeeze(0)
|
||||||
o.add_(ps * mask)
|
o.add_(ps * mask)
|
||||||
o_d.add_(mask)
|
o_d.add_(mask)
|
||||||
|
|
||||||
if pbar is not None:
|
if pbar is not None:
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
output[b:b+1] = out/out_div
|
# Fallback to non-tiled if NaN
|
||||||
return output
|
if torch.isnan(out).any():
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.warning("NaN detected in tiled output, falling back to non-tiled")
|
||||||
|
s_fp32 = s.to(dtype=torch.float32)
|
||||||
|
result = fp32_function(s_fp32).to(output_device, dtype=torch.float32)
|
||||||
|
if result.shape == output_shape[1:]:
|
||||||
|
output[b:b+1] = result
|
||||||
|
else:
|
||||||
|
result = result.narrow(1, 0, output_shape[1])
|
||||||
|
for d in range(dims):
|
||||||
|
result = result.narrow(d + 2, 0, output_shape[d + 2])
|
||||||
|
output[b:b+1] = result
|
||||||
|
else:
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"out stats: min={out.min():.4f}, max={out.max():.4f}")
|
||||||
|
logging.debug(f"out_div stats: min={out_div.min():.4f}, max={out_div.max():.4f}")
|
||||||
|
output[b:b+1] = out / out_div
|
||||||
|
|
||||||
def tiled_scale(samples, function, tile_x=64, tile_y=64, overlap = 8, upscale_amount = 4, out_channels = 3, output_device="cpu", pbar = None):
|
if pbar is not None:
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
def tiled_scale(samples, function, tile_x=64, tile_y=64, overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", pbar=None):
|
||||||
|
"""Wrapper for 2D tiled scaling."""
|
||||||
return tiled_scale_multidim(samples, function, (tile_y, tile_x), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=output_device, pbar=pbar)
|
return tiled_scale_multidim(samples, function, (tile_y, tile_x), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=output_device, pbar=pbar)
|
||||||
|
|
||||||
PROGRESS_BAR_ENABLED = True
|
PROGRESS_BAR_ENABLED = True
|
||||||
|
|||||||
@ -8,6 +8,7 @@ from comfy.model_management import get_torch_device, vae_dtype, soft_empty_cache
|
|||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
import latent_preview
|
import latent_preview
|
||||||
import logging
|
import logging
|
||||||
|
import traceback
|
||||||
|
|
||||||
# Global flag for profiling
|
# Global flag for profiling
|
||||||
PROFILING_ENABLED = args.profile
|
PROFILING_ENABLED = args.profile
|
||||||
@ -41,15 +42,18 @@ def profile_cuda_sync(is_gpu, message="CUDA sync"):
|
|||||||
logging.debug(f"{message} took {time.time() - sync_start:.3f} s")
|
logging.debug(f"{message} took {time.time() - sync_start:.3f} s")
|
||||||
|
|
||||||
def is_fp16_safe(device):
|
def is_fp16_safe(device):
|
||||||
"""Check if FP16 is safe for the GPU (disabled for GTX 1660/Turing)."""
|
"""Check if FP16 is safe for the GPU (disabled for Turing)."""
|
||||||
if device.type != 'cuda':
|
if device.type != 'cuda':
|
||||||
return False
|
return False
|
||||||
if device in _fp16_safe_cache:
|
if device in _fp16_safe_cache:
|
||||||
return _fp16_safe_cache[device]
|
return _fp16_safe_cache[device]
|
||||||
try:
|
try:
|
||||||
props = torch.cuda.get_device_properties(device)
|
props = torch.cuda.get_device_properties(device)
|
||||||
is_safe = props.major >= 8 or props.compute_capability[0] > 7
|
# Disable FP16 for Turing (major == 7) and earlier architectures
|
||||||
|
is_safe = props.major >= 8 # Allow FP16 only for Ampere (8.x) and later
|
||||||
_fp16_safe_cache[device] = is_safe
|
_fp16_safe_cache[device] = is_safe
|
||||||
|
if DEBUG_ENABLED:
|
||||||
|
logging.debug(f"FP16 safety check for {props.name}: major={props.major}, is_safe={is_safe}")
|
||||||
return is_safe
|
return is_safe
|
||||||
except Exception:
|
except Exception:
|
||||||
_fp16_safe_cache[device] = False
|
_fp16_safe_cache[device] = False
|
||||||
@ -72,7 +76,8 @@ def clear_vram(device, threshold=0.5, min_free=1.5):
|
|||||||
mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3
|
mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3
|
||||||
critical_threshold = 0.05 * mem_total + 0.1 # 5% VRAM + 100 MB
|
critical_threshold = 0.05 * mem_total + 0.1 # 5% VRAM + 100 MB
|
||||||
if mem_allocated > threshold * mem_total or (mem_total - mem_allocated) < max(min_free, critical_threshold):
|
if mem_allocated > threshold * mem_total or (mem_total - mem_allocated) < max(min_free, critical_threshold):
|
||||||
logging.debug(f"Clearing VRAM: allocated {mem_allocated:.2f} GB, free {mem_total - mem_allocated:.2f} GB, threshold {critical_threshold:.2f} GB")
|
if PROFILING_ENABLED:
|
||||||
|
logging.debug(f"Clearing VRAM: allocated {mem_allocated:.2f} GB, free {mem_total - mem_allocated:.2f} GB, threshold {critical_threshold:.2f} GB")
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
#soft_empty_cache(clear=False)
|
#soft_empty_cache(clear=False)
|
||||||
mem_after = torch.cuda.memory_allocated(device) / 1024**3
|
mem_after = torch.cuda.memory_allocated(device) / 1024**3
|
||||||
@ -362,9 +367,9 @@ def fast_vae_decode(vae, samples):
|
|||||||
def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size=64, temporal_overlap=8):
|
def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size=64, temporal_overlap=8):
|
||||||
"""Fast VAE decoding with tiling for low VRAM, consistent with fast_vae_decode."""
|
"""Fast VAE decoding with tiling for low VRAM, consistent with fast_vae_decode."""
|
||||||
device, dtype, is_gpu = initialize_device_and_dtype(vae)
|
device, dtype, is_gpu = initialize_device_and_dtype(vae)
|
||||||
vae_dtype = vae_dtype(device=device)
|
vae_dtype_val = vae_dtype(device=device)
|
||||||
if DEBUG_ENABLED:
|
if DEBUG_ENABLED:
|
||||||
logging.debug(f"VAE dtype: {vae_dtype}")
|
logging.debug(f"VAE dtype: {vae_dtype_val}")
|
||||||
logging.debug(f"Pre-VAE checkpoint: {time.time()}")
|
logging.debug(f"Pre-VAE checkpoint: {time.time()}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@ -378,7 +383,7 @@ def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size
|
|||||||
mem_allocated = torch.cuda.memory_allocated(device) / 1024**3
|
mem_allocated = torch.cuda.memory_allocated(device) / 1024**3
|
||||||
free_mem = mem_total - mem_allocated
|
free_mem = mem_total - mem_allocated
|
||||||
# Estimate memory for tiled decoding (conservative, ~50% of full decode)
|
# Estimate memory for tiled decoding (conservative, ~50% of full decode)
|
||||||
vae_memory_required = (vae.memory_used_decode(samples["samples"].shape, vae_dtype) / 1024**3 * 0.5
|
vae_memory_required = (vae.memory_used_decode(samples["samples"].shape, vae_dtype_val) / 1024**3 * 0.5
|
||||||
if hasattr(vae, 'memory_used_decode') else 0.75)
|
if hasattr(vae, 'memory_used_decode') else 0.75)
|
||||||
if PROFILING_ENABLED:
|
if PROFILING_ENABLED:
|
||||||
logging.debug(f"VRAM before tiled VAE: {mem_allocated:.2f} GB / {mem_total:.2f} GB")
|
logging.debug(f"VRAM before tiled VAE: {mem_allocated:.2f} GB / {mem_total:.2f} GB")
|
||||||
@ -408,7 +413,7 @@ def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size
|
|||||||
latent_samples = samples["samples"]
|
latent_samples = samples["samples"]
|
||||||
if PROFILING_ENABLED:
|
if PROFILING_ENABLED:
|
||||||
logging.debug(f"Latent samples device: {latent_samples.device}, dtype: {latent_samples.dtype}")
|
logging.debug(f"Latent samples device: {latent_samples.device}, dtype: {latent_samples.dtype}")
|
||||||
latent_samples = optimized_transfer(latent_samples, device, vae_dtype)
|
latent_samples = optimized_transfer(latent_samples, device, vae_dtype_val)
|
||||||
if is_gpu and force_channels_last():
|
if is_gpu and force_channels_last():
|
||||||
latent_samples = latent_samples.to(memory_format=torch.channels_last)
|
latent_samples = latent_samples.to(memory_format=torch.channels_last)
|
||||||
vae.first_stage_model.to(memory_format=torch.channels_last)
|
vae.first_stage_model.to(memory_format=torch.channels_last)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user