support older torch

This commit is contained in:
kijai
2024-08-01 11:16:39 +03:00
parent 6941b9001a
commit 9c7322184f
3 changed files with 46 additions and 30 deletions
+37 -26
View File
@@ -9,22 +9,30 @@ import warnings
from functools import partial
from typing import Tuple, Type
import torch
import torch.nn.functional as F
from torch import nn, Tensor
from ....sam2.modeling.position_encoding import apply_rotary_enc, compute_axial_cis
from ....sam2.modeling.sam2_utils import MLP
from ....sam2.utils.misc import get_sdpa_settings
OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
try:
from torch.nn.attention import SDPBackend, sdpa_kernel
except: #old torch
from torch.nn.functional import scaled_dot_product_attention as sdpa_kernel
from torch._C import _SDPBackend as SDPBackend
backends = []
if USE_FLASH_ATTN:
backends.append(SDPBackend.FLASH_ATTENTION)
if MATH_KERNEL_ON:
backends.append(SDPBackend.MATH)
if OLD_GPU:
backends.append(SDPBackend.EFFICIENT_ATTENTION)
OLD_TORCH = False
except:
OLD_TORCH = True
warnings.simplefilter(action="ignore", category=FutureWarning)
OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
class TwoWayTransformer(nn.Module):
def __init__(
@@ -250,17 +258,18 @@ class Attention(nn.Module):
dropout_p = self.dropout_p if self.training else 0.0
# Attention
# Determine backends to enable
backends = [SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION]
if USE_FLASH_ATTN:
backends.append(SDPBackend.FLASH_ATTENTION)
if OLD_GPU and dropout_p > 0.0:
backends.append(SDPBackend.MATH)
if OLD_GPU:
backends.append(SDPBackend.EFFICIENT_ATTENTION)
with sdpa_kernel(backends):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
if not OLD_TORCH:
if not MATH_KERNEL_ON and OLD_GPU and dropout_p > 0.0:
backends.append(SDPBackend.MATH)
with sdpa_kernel(backends):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
else:
with torch.backends.cuda.sdp_kernel(
enable_flash=USE_FLASH_ATTN,
enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
enable_mem_efficient=OLD_GPU,
):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
out = self._recombine_heads(out)
out = self.out_proj(out)
@@ -320,16 +329,18 @@ class RoPEAttention(Attention):
dropout_p = self.dropout_p if self.training else 0.0
# Attention
backends = [SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION]
if USE_FLASH_ATTN:
backends.append(SDPBackend.FLASH_ATTENTION)
if OLD_GPU and dropout_p > 0.0:
backends.append(SDPBackend.MATH)
if OLD_GPU:
backends.append(SDPBackend.EFFICIENT_ATTENTION)
with sdpa_kernel(backends):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
if not OLD_TORCH:
if not MATH_KERNEL_ON and OLD_GPU and dropout_p > 0.0:
backends.append(SDPBackend.MATH))
with sdpa_kernel(backends):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
else:
with torch.backends.cuda.sdp_kernel(
enable_flash=USE_FLASH_ATTN,
enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
enable_mem_efficient=OLD_GPU,
):
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
out = self._recombine_heads(out)
out = self.out_proj(out)
+7 -2
View File
@@ -7,7 +7,7 @@
from collections import OrderedDict
import torch
import numpy as np
from tqdm import tqdm
from ..sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
@@ -41,7 +41,7 @@ class SAM2VideoPredictor(SAM2Base):
images,
video_height,
video_width,
device,
device='cuda',
offload_video_to_cpu=False,
offload_state_to_cpu=False,
async_loading_frames=False,
@@ -165,8 +165,13 @@ class SAM2VideoPredictor(SAM2Base):
mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
if not isinstance(points, torch.Tensor):
if isinstance(points, list) and all(isinstance(p, np.ndarray) for p in points):
points = np.array(points)
points = torch.tensor(points, dtype=torch.float32)
if not isinstance(labels, torch.Tensor):
if isinstance(labels, list) and all(isinstance(l, np.ndarray) for l in labels):
labels = np.array(labels)
labels = torch.tensor(labels, dtype=torch.int32)
if points.dim() == 2:
points = points.unsqueeze(0) # add batch dimension
+2 -2
View File
@@ -12,13 +12,13 @@ import numpy as np
import torch
from PIL import Image
from tqdm import tqdm
import platform
def get_sdpa_settings():
if torch.cuda.is_available():
old_gpu = torch.cuda.get_device_properties(0).major < 7
# only use Flash Attention on Ampere (8.0) or newer GPUs
use_flash_attn = torch.cuda.get_device_properties(0).major >= 8
use_flash_attn = torch.cuda.get_device_properties(0).major >= 8 and platform.system() == 'Linux'
if not use_flash_attn:
warnings.warn(
"Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.",