From 0e4b8df789d9b78a0804711a9ea64e6426224233 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=B0=A2=E7=BF=8A=E5=87=A1?= Date: Fri, 21 Aug 2026 17:47:40 +0800 Subject: [PATCH] Migrate torch.cuda.amp.autocast to torch.amp torch.cuda.amp.autocast is deprecated since torch 2.4 and scheduled for removal. Six imports in the core flagai/ package switch to a try/except import (torch.amp on torch >= 2.0, torch.cuda.amp below, which is not yet deprecated), keeping the documented torch >= 1.8 floor intact. Call sites gain the required device arg: Predictor, AltDiffusion/AltDiffusionM18, Unet, openaimodel, and the autograd function's **kwargs form in diffusionmodules/util.py (via a version-guarded helper). Fixes #592 Co-Authored-By: Claude --- flagai/model/mm/AltDiffusion.py | 15 +++++++++++++-- flagai/model/mm/AltDiffusionM18.py | 16 +++++++++++++--- flagai/model/mm/Unets/Unet.py | 15 +++++++++++++-- .../mm/modules/diffusionmodules/openaimodel.py | 15 +++++++++++++-- flagai/model/mm/modules/diffusionmodules/util.py | 8 +++++++- flagai/model/predictor/predictor.py | 15 +++++++++++++-- 6 files changed, 72 insertions(+), 12 deletions(-) diff --git a/flagai/model/mm/AltDiffusion.py b/flagai/model/mm/AltDiffusion.py index 62c8ae60..71ff7f92 100755 --- a/flagai/model/mm/AltDiffusion.py +++ b/flagai/model/mm/AltDiffusion.py @@ -14,7 +14,18 @@ from flagai.model.mm.utils import make_beta_schedule, extract_into_tensor, noise_like from flagai.model.mm.Sampler import DDIMSampler from flagai.model.base_model import BaseModel -from torch.cuda.amp import autocast as autocast +try: # torch >= 2.0 + from torch.amp import autocast +except ImportError: # torch < 2.0 + from torch.cuda.amp import autocast + +def _autocast(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + + __conditioning_keys__ = { 'concat': 'c_concat', @@ -1266,7 +1277,7 @@ def apply_model(self, x_noisy, t, cond, return_ids=False): x_recon = fold(o) / normalization else: - with autocast(): + with _autocast(): x_recon = self.model(x_noisy, t, **cond) if isinstance(x_recon, tuple) and not return_ids: diff --git a/flagai/model/mm/AltDiffusionM18.py b/flagai/model/mm/AltDiffusionM18.py index acb30e95..01fb4eb9 100755 --- a/flagai/model/mm/AltDiffusionM18.py +++ b/flagai/model/mm/AltDiffusionM18.py @@ -14,9 +14,19 @@ from flagai.model.mm.utils import make_beta_schedule, extract_into_tensor, noise_like from flagai.model.mm.Sampler import DDIMSampler from flagai.model.base_model import BaseModel -from torch.cuda.amp import autocast as autocast +try: # torch >= 2.0 + from torch.amp import autocast +except ImportError: # torch < 2.0 + from torch.cuda.amp import autocast + +def _autocast(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + + import pytorch_lightning as pl -from torch.cuda.amp import autocast as autocast __conditioning_keys__ = { @@ -891,7 +901,7 @@ def apply_model(self, x_noisy, t, cond, return_ids=False): cond = [cond] key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn' cond = {key: cond} - with autocast(): + with _autocast(): x_recon = self.model(x_noisy, t, **cond) if isinstance(x_recon, tuple) and not return_ids: diff --git a/flagai/model/mm/Unets/Unet.py b/flagai/model/mm/Unets/Unet.py index 798f8ff5..3d604e38 100644 --- a/flagai/model/mm/Unets/Unet.py +++ b/flagai/model/mm/Unets/Unet.py @@ -16,7 +16,18 @@ timestep_embedding, ) from flagai.model.mm.attentions.attention import SpatialTransformer -from torch.cuda.amp import autocast as autocast +try: # torch >= 2.0 + from torch.amp import autocast +except ImportError: # torch < 2.0 + from torch.cuda.amp import autocast + +def _autocast(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + + # dummy replace def convert_module_to_f16(x): @@ -78,7 +89,7 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock): def forward(self, x, emb, context=None): for layer in self: if isinstance(layer, TimestepBlock): - with autocast(): + with _autocast(): x = layer(x, emb) elif isinstance(layer, SpatialTransformer): x = layer(x, context) diff --git a/flagai/model/mm/modules/diffusionmodules/openaimodel.py b/flagai/model/mm/modules/diffusionmodules/openaimodel.py index 4c85019b..4317bb75 100644 --- a/flagai/model/mm/modules/diffusionmodules/openaimodel.py +++ b/flagai/model/mm/modules/diffusionmodules/openaimodel.py @@ -7,7 +7,18 @@ import torch.nn as nn import torch.nn.functional as F from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ -from torch.cuda.amp import autocast as autocast +try: # torch >= 2.0 + from torch.amp import autocast +except ImportError: # torch < 2.0 + from torch.cuda.amp import autocast + +def _autocast(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + + from flagai.model.mm.modules.diffusionmodules.util import ( checkpoint, @@ -82,7 +93,7 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock): def forward(self, x, emb, context=None, heypernetwork=None): for layer in self: if isinstance(layer, TimestepBlock): - with autocast(): + with _autocast(): x = layer(x, emb) elif isinstance(layer, SpatialTransformer): x = layer(x, context) diff --git a/flagai/model/mm/modules/diffusionmodules/util.py b/flagai/model/mm/modules/diffusionmodules/util.py index 5b09d609..51872ca8 100644 --- a/flagai/model/mm/modules/diffusionmodules/util.py +++ b/flagai/model/mm/modules/diffusionmodules/util.py @@ -17,6 +17,12 @@ from flagai.model.mm.utils import instantiate_from_config +def _autocast_ctx(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3): if schedule == "linear": @@ -133,7 +139,7 @@ def forward(ctx, run_function, length, *args): def backward(ctx, *output_grads): ctx.input_tensors = [x.detach().requires_grad_(True) for x in ctx.input_tensors] with torch.enable_grad(), \ - torch.cuda.amp.autocast(**ctx.gpu_autocast_kwargs): + _autocast_ctx(**ctx.gpu_autocast_kwargs): # Fixes a bug where the first op in run_function modifies the # Tensor storage in place, which is not allowed for detach()'d # Tensors. diff --git a/flagai/model/predictor/predictor.py b/flagai/model/predictor/predictor.py index fc41c8b2..92c78967 100755 --- a/flagai/model/predictor/predictor.py +++ b/flagai/model/predictor/predictor.py @@ -18,7 +18,18 @@ import time from contextlib import contextmanager, nullcontext from einops import rearrange -from torch.cuda.amp import autocast as autocast +try: # torch >= 2.0 + from torch.amp import autocast +except ImportError: # torch < 2.0 + from torch.cuda.amp import autocast + +def _autocast(**kwargs): + """torch.amp.autocast("cuda") on torch >= 2.0, else torch.cuda.amp.autocast.""" + if hasattr(torch, "amp"): + return torch.amp.autocast("cuda", **kwargs) + return torch.cuda.amp.autocast(**kwargs) + + from .aquila import aquila_generate class Predictor: @@ -456,7 +467,7 @@ def predict_generate_images(self, unconditional_conditioning=uc, eta=ddim_eta, x_T=start_code) - with autocast(): + with _autocast(): x_samples_ddim = self.model.decode_first_stage( samples_ddim) x_samples_ddim = torch.clamp(