Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 13 additions & 2 deletions flagai/model/mm/AltDiffusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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',
Expand Down Expand Up @@ -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:
Expand Down
16 changes: 13 additions & 3 deletions flagai/model/mm/AltDiffusionM18.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__ = {
Expand Down Expand Up @@ -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:
Expand Down
15 changes: 13 additions & 2 deletions flagai/model/mm/Unets/Unet.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand Down
15 changes: 13 additions & 2 deletions flagai/model/mm/modules/diffusionmodules/openaimodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
8 changes: 7 additions & 1 deletion flagai/model/mm/modules/diffusionmodules/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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":
Expand Down Expand Up @@ -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.
Expand Down
15 changes: 13 additions & 2 deletions flagai/model/predictor/predictor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down