refactor: Remove dead code from utils/common.py (#20668)
This commit is contained in:
@@ -85,7 +85,6 @@ from PIL import Image
|
||||
from starlette.routing import Mount
|
||||
from torch import nn
|
||||
from torch.library import Library
|
||||
from torch.profiler import ProfilerActivity, profile, record_function
|
||||
from torch.utils._contextlib import _DecoratorContextManager
|
||||
from typing_extensions import Literal
|
||||
|
||||
@@ -363,17 +362,6 @@ def get_int_env_var(name: str, default: int = 0) -> int:
|
||||
return default
|
||||
|
||||
|
||||
def get_float_env_var(name: str, default: float = 0.0) -> float:
|
||||
# FIXME: move your environment variable to sglang.srt.environ
|
||||
value = os.getenv(name)
|
||||
if value is None or not value.strip():
|
||||
return default
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
|
||||
def support_triton(backend: str) -> bool:
|
||||
return backend not in ["torch_native", "intel_amx"]
|
||||
|
||||
@@ -700,85 +688,6 @@ def set_random_seed(seed: int) -> None:
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
|
||||
def decode_video_base64(video_base64):
|
||||
from PIL import Image
|
||||
|
||||
# Decode the base64 string
|
||||
video_bytes = pybase64.b64decode(video_base64, validate=True)
|
||||
|
||||
# Placeholder for the start indices of each PNG image
|
||||
img_starts = []
|
||||
|
||||
frame_format = "PNG" # str(os.getenv('FRAME_FORMAT', "JPEG"))
|
||||
|
||||
assert frame_format in [
|
||||
"PNG",
|
||||
"JPEG",
|
||||
], "FRAME_FORMAT must be either 'PNG' or 'JPEG'"
|
||||
|
||||
if frame_format == "PNG":
|
||||
# Find each PNG start signature to isolate images
|
||||
i = 0
|
||||
while i < len(video_bytes) - 7: # Adjusted for the length of the PNG signature
|
||||
# Check if we found the start of a PNG file
|
||||
if (
|
||||
video_bytes[i] == 0x89
|
||||
and video_bytes[i + 1] == 0x50
|
||||
and video_bytes[i + 2] == 0x4E
|
||||
and video_bytes[i + 3] == 0x47
|
||||
and video_bytes[i + 4] == 0x0D
|
||||
and video_bytes[i + 5] == 0x0A
|
||||
and video_bytes[i + 6] == 0x1A
|
||||
and video_bytes[i + 7] == 0x0A
|
||||
):
|
||||
img_starts.append(i)
|
||||
i += 8 # Skip the PNG signature
|
||||
else:
|
||||
i += 1
|
||||
else:
|
||||
# Find each JPEG start (0xFFD8) to isolate images
|
||||
i = 0
|
||||
while (
|
||||
i < len(video_bytes) - 1
|
||||
): # Adjusted for the length of the JPEG SOI signature
|
||||
# Check if we found the start of a JPEG file
|
||||
if video_bytes[i] == 0xFF and video_bytes[i + 1] == 0xD8:
|
||||
img_starts.append(i)
|
||||
# Move to the next byte to continue searching for the next image start
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
|
||||
frames = []
|
||||
for start_idx in img_starts:
|
||||
# Assuming each image is back-to-back, the end of one image is the start of another
|
||||
# The last image goes until the end of the byte string
|
||||
end_idx = (
|
||||
img_starts[img_starts.index(start_idx) + 1]
|
||||
if img_starts.index(start_idx) + 1 < len(img_starts)
|
||||
else len(video_bytes)
|
||||
)
|
||||
img_bytes = video_bytes[start_idx:end_idx]
|
||||
|
||||
# Convert bytes to a PIL Image
|
||||
img = Image.open(BytesIO(img_bytes))
|
||||
|
||||
# Convert PIL Image to a NumPy array
|
||||
frame = np.array(img)
|
||||
|
||||
# Append the frame to the list of frames
|
||||
frames.append(frame)
|
||||
|
||||
# Ensure there's at least one frame to avoid errors with np.stack
|
||||
if frames:
|
||||
return np.stack(frames, axis=0), img.size
|
||||
else:
|
||||
return np.array([]), (
|
||||
0,
|
||||
0,
|
||||
) # Return an empty array and size tuple if no frames were found
|
||||
|
||||
|
||||
def load_audio(
|
||||
audio_file: str, sr: Optional[int] = None, mono: bool = True
|
||||
) -> np.ndarray:
|
||||
@@ -1341,70 +1250,6 @@ def point_to_point_pyobj(
|
||||
return []
|
||||
|
||||
|
||||
step_counter = 0
|
||||
|
||||
|
||||
def pytorch_profile(name, func, *args, data_size=-1):
|
||||
"""
|
||||
Args:
|
||||
name (string): the name of recorded function.
|
||||
func: the function to be profiled.
|
||||
args: the arguments of the profiled function.
|
||||
data_size (int): some measurement of the computation complexity.
|
||||
Usually, it could be the batch size.
|
||||
"""
|
||||
global step_counter
|
||||
os.makedirs("trace", exist_ok=True)
|
||||
with profile(
|
||||
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
|
||||
# schedule=torch.profiler.schedule(wait=1, warmup=1, active=3, repeat=2),
|
||||
# on_trace_ready=tensorboard_trace_handler('./log_dir'),
|
||||
record_shapes=True,
|
||||
profile_memory=True,
|
||||
with_stack=True,
|
||||
) as prof:
|
||||
with record_function(name):
|
||||
with open(f"trace/size_{step_counter}.json", "w") as f:
|
||||
json.dump({"size": data_size}, f)
|
||||
result = func(*args)
|
||||
prof.export_chrome_trace(f"trace/{name}_{step_counter}.json")
|
||||
step_counter += 1
|
||||
return result
|
||||
|
||||
|
||||
def dump_to_file(dirpath, name, value):
|
||||
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
||||
|
||||
if get_tensor_model_parallel_rank() != 0:
|
||||
return
|
||||
|
||||
os.makedirs(dirpath, exist_ok=True)
|
||||
if value.dtype is torch.bfloat16:
|
||||
value = value.float()
|
||||
value = value.cpu().numpy()
|
||||
output_filename = os.path.join(dirpath, f"pytorch_dump_{name}.npy")
|
||||
logger.info(f"Dump a tensor to {output_filename}. Shape = {value.shape}")
|
||||
np.save(output_filename, value)
|
||||
|
||||
|
||||
def is_triton_3():
|
||||
return triton.__version__.startswith("3.")
|
||||
|
||||
|
||||
def maybe_torch_compile(*args, **kwargs):
|
||||
"""
|
||||
torch.compile does not work for triton 2.2.0, which is needed in xlm1's jax.
|
||||
Therefore, we disable it here.
|
||||
"""
|
||||
|
||||
def decorator(func):
|
||||
if is_triton_3():
|
||||
return torch.compile(*args, **kwargs)(func)
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def delete_directory(dirpath):
|
||||
try:
|
||||
# This will remove the directory and all its contents
|
||||
|
||||
Reference in New Issue
Block a user