Skip to content
Closed
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
8 changes: 8 additions & 0 deletions torchbenchmark/models/sam_fast/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,14 @@
import cv2
import numpy as np
import torch
import triton.language as tl

# The optional flash_4 kernel requires Triton's legacy block-pointer frontend.
# Disable it before importing segment_anything_fast so the model uses its SDPA
# fallback when those APIs are unavailable.
if not hasattr(tl, "make_block_ptr") or not hasattr(tl, "advance"):
os.environ["SEGMENT_ANYTHING_FAST_USE_FLASH_4"] = "0"

from segment_anything_fast.build_sam import sam_model_fast_registry
from segment_anything_fast.predictor import SamPredictor
from torchbenchmark import DATA_PATH
Expand Down
Loading