From bd34f338ac505ea79e43968753968a464060e609 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:55:08 -0700 Subject: [PATCH] Fix float64 device in ltx diffusion decoder. (#15516) --- comfy/ldm/lightricks/vae/na_diffusion_decoder.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py index 8a172e101..ec539c665 100644 --- a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py +++ b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py @@ -22,6 +22,7 @@ import torch import torch.nn.functional as F from einops import rearrange from torch import nn +import comfy.model_management from comfy.ldm.lightricks.model import get_timestep_embedding from .causal_video_autoencoder import Encoder, processor @@ -74,8 +75,12 @@ def default_rope_dim_split(head_dim): def rope_inv_freqs(dim, base=10000.0, device=None): + out_device = device + if not comfy.model_management.supports_fp64(device): + device = torch.device("cpu") + exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim - return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(torch.float32) + return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(dtype=torch.float32, device=out_device) def _rope_tables(lengths, inv_freqs, device):