Tiny move files to utils folder (#11166)
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
# Temporarily do this to avoid changing all imports in the repo
|
||||
from .common import *
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,448 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Utilities for Huggingface Transformers."""
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Type, Union
|
||||
|
||||
import torch
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
AutoProcessor,
|
||||
AutoTokenizer,
|
||||
GenerationConfig,
|
||||
PretrainedConfig,
|
||||
PreTrainedTokenizer,
|
||||
PreTrainedTokenizerBase,
|
||||
PreTrainedTokenizerFast,
|
||||
)
|
||||
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
|
||||
|
||||
from sglang.srt.configs import (
|
||||
ChatGLMConfig,
|
||||
DbrxConfig,
|
||||
DeepseekVL2Config,
|
||||
DotsOCRConfig,
|
||||
DotsVLMConfig,
|
||||
ExaoneConfig,
|
||||
FalconH1Config,
|
||||
KimiVLConfig,
|
||||
LongcatFlashConfig,
|
||||
MultiModalityConfig,
|
||||
Qwen3NextConfig,
|
||||
Step3VLConfig,
|
||||
)
|
||||
from sglang.srt.configs.internvl import InternVLChatConfig
|
||||
from sglang.srt.connector import create_remote_connector
|
||||
from sglang.srt.utils import is_remote_url, logger, lru_cache_frozenset
|
||||
|
||||
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
ChatGLMConfig.model_type: ChatGLMConfig,
|
||||
DbrxConfig.model_type: DbrxConfig,
|
||||
ExaoneConfig.model_type: ExaoneConfig,
|
||||
DeepseekVL2Config.model_type: DeepseekVL2Config,
|
||||
MultiModalityConfig.model_type: MultiModalityConfig,
|
||||
KimiVLConfig.model_type: KimiVLConfig,
|
||||
InternVLChatConfig.model_type: InternVLChatConfig,
|
||||
Step3VLConfig.model_type: Step3VLConfig,
|
||||
LongcatFlashConfig.model_type: LongcatFlashConfig,
|
||||
Qwen3NextConfig.model_type: Qwen3NextConfig,
|
||||
FalconH1Config.model_type: FalconH1Config,
|
||||
DotsVLMConfig.model_type: DotsVLMConfig,
|
||||
DotsOCRConfig.model_type: DotsOCRConfig,
|
||||
}
|
||||
|
||||
for name, cls in _CONFIG_REGISTRY.items():
|
||||
with contextlib.suppress(ValueError):
|
||||
AutoConfig.register(name, cls)
|
||||
|
||||
|
||||
def download_from_hf(
|
||||
model_path: str,
|
||||
allow_patterns: Optional[Union[str, list]] = None,
|
||||
):
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
|
||||
if not allow_patterns:
|
||||
allow_patterns = ["*.json", "*.bin", "*.model"]
|
||||
|
||||
return snapshot_download(model_path, allow_patterns=allow_patterns)
|
||||
|
||||
|
||||
def get_hf_text_config(config: PretrainedConfig):
|
||||
"""Get the "sub" config relevant to llm for multi modal models.
|
||||
No op for pure text models.
|
||||
"""
|
||||
if config.architectures is not None:
|
||||
class_name = config.architectures[0]
|
||||
if class_name.startswith("Llava") and class_name.endswith("ForCausalLM"):
|
||||
# We support non-hf version of llava models, so we do not want to
|
||||
# read the wrong values from the unused default text_config.
|
||||
# NOTE(HandH1998): We set `torch_dtype` of config to `torch.float16` for the weights, as
|
||||
# `torch.float16` is default used for image features in `python/sglang/srt/models/llava.py`.
|
||||
setattr(config, "torch_dtype", torch.float16)
|
||||
return config
|
||||
|
||||
if hasattr(config, "text_config"):
|
||||
# The code operates under the assumption that text_config should have
|
||||
# `num_attention_heads` (among others). Assert here to fail early
|
||||
# if transformers config doesn't align with this assumption.
|
||||
assert hasattr(config.text_config, "num_attention_heads")
|
||||
return config.text_config
|
||||
if hasattr(config, "language_config"):
|
||||
return config.language_config
|
||||
if hasattr(config, "thinker_config"):
|
||||
# qwen2.5 omni
|
||||
thinker_config = config.thinker_config
|
||||
if hasattr(thinker_config, "text_config"):
|
||||
setattr(
|
||||
thinker_config.text_config,
|
||||
"torch_dtype",
|
||||
getattr(thinker_config, "torch_dtype", None),
|
||||
)
|
||||
return thinker_config.text_config
|
||||
return thinker_config
|
||||
else:
|
||||
return config
|
||||
|
||||
|
||||
@lru_cache_frozenset(maxsize=32)
|
||||
def get_config(
|
||||
model: str,
|
||||
trust_remote_code: bool,
|
||||
revision: Optional[str] = None,
|
||||
model_override_args: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
if is_gguf:
|
||||
kwargs["gguf_file"] = model
|
||||
model = Path(model).parent
|
||||
|
||||
if is_remote_url(model):
|
||||
# BaseConnector implements __del__() to clean up the local dir.
|
||||
# Since config files need to exist all the time, so we DO NOT use
|
||||
# with statement to avoid closing the client.
|
||||
client = create_remote_connector(model)
|
||||
client.pull_files(ignore_pattern=["*.pt", "*.safetensors", "*.bin"])
|
||||
model = client.get_local_dir()
|
||||
|
||||
config = AutoConfig.from_pretrained(
|
||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||
)
|
||||
if (
|
||||
config.architectures is not None
|
||||
and config.architectures[0] == "Phi4MMForCausalLM"
|
||||
):
|
||||
# Phi4MMForCausalLM uses a hard-coded vision_config. See:
|
||||
# https://github.com/vllm-project/vllm/blob/6071e989df1531b59ef35568f83f7351afb0b51e/vllm/model_executor/models/phi4mm.py#L71
|
||||
# We set it here to support cases where num_attention_heads is not divisible by the TP size.
|
||||
from transformers import SiglipVisionConfig
|
||||
|
||||
vision_config = {
|
||||
"hidden_size": 1152,
|
||||
"image_size": 448,
|
||||
"intermediate_size": 4304,
|
||||
"model_type": "siglip_vision_model",
|
||||
"num_attention_heads": 16,
|
||||
"num_hidden_layers": 26, # Model is originally 27-layer, we only need the first 26 layers for feature extraction.
|
||||
"patch_size": 14,
|
||||
}
|
||||
config.vision_config = SiglipVisionConfig(**vision_config)
|
||||
text_config = get_hf_text_config(config=config)
|
||||
|
||||
if isinstance(model, str) and text_config is not None:
|
||||
for key, val in text_config.__dict__.items():
|
||||
if not hasattr(config, key) and getattr(text_config, key, None) is not None:
|
||||
setattr(config, key, val)
|
||||
|
||||
if config.model_type in _CONFIG_REGISTRY:
|
||||
config_class = _CONFIG_REGISTRY[config.model_type]
|
||||
config = config_class.from_pretrained(model, revision=revision)
|
||||
# NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`.
|
||||
setattr(config, "_name_or_path", model)
|
||||
|
||||
if isinstance(model, str) and config.model_type == "internvl_chat":
|
||||
for key, val in config.llm_config.__dict__.items():
|
||||
if not hasattr(config, key):
|
||||
setattr(config, key, val)
|
||||
|
||||
if config.model_type == "multi_modality":
|
||||
config.update({"architectures": ["MultiModalityCausalLM"]})
|
||||
|
||||
if model_override_args:
|
||||
config.update(model_override_args)
|
||||
|
||||
# Special architecture mapping check for GGUF models
|
||||
if is_gguf:
|
||||
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
|
||||
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
|
||||
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
|
||||
config.update({"architectures": [model_type]})
|
||||
|
||||
return config
|
||||
|
||||
|
||||
@lru_cache_frozenset(maxsize=32)
|
||||
def get_generation_config(
|
||||
model: str,
|
||||
trust_remote_code: bool,
|
||||
revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
try:
|
||||
return GenerationConfig.from_pretrained(
|
||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||
)
|
||||
except OSError as e:
|
||||
return None
|
||||
|
||||
|
||||
# Qwen-1M related
|
||||
def get_sparse_attention_config(
|
||||
model: str,
|
||||
sparse_attention_config_filename: str = "sparse_attention_config.json",
|
||||
) -> Dict[str, Any]:
|
||||
is_local = os.path.isdir(model)
|
||||
if not is_local:
|
||||
# Download the config files.
|
||||
model = download_from_hf(model, allow_patterns=["*.json"])
|
||||
|
||||
config_file = os.path.join(model, sparse_attention_config_filename)
|
||||
if not os.path.exists(config_file):
|
||||
return {}
|
||||
|
||||
# Load the sparse attention config.
|
||||
with open(config_file) as f:
|
||||
config = json.load(f)
|
||||
return config
|
||||
|
||||
|
||||
# Models don't use the same configuration key for determining the maximum
|
||||
# context length. Store them here so we can sanely check them.
|
||||
# NOTE: The ordering here is important. Some models have two of these and we
|
||||
# have a preference for which value gets used.
|
||||
CONTEXT_LENGTH_KEYS = [
|
||||
"max_sequence_length",
|
||||
"seq_length",
|
||||
"max_seq_len",
|
||||
"model_max_length",
|
||||
"max_position_embeddings",
|
||||
]
|
||||
|
||||
|
||||
def get_context_length(config):
|
||||
"""Get the context length of a model from a huggingface model configs."""
|
||||
text_config = config
|
||||
rope_scaling = getattr(text_config, "rope_scaling", None)
|
||||
if rope_scaling:
|
||||
rope_scaling_factor = rope_scaling.get("factor", 1)
|
||||
if "original_max_position_embeddings" in rope_scaling:
|
||||
rope_scaling_factor = 1
|
||||
if rope_scaling.get("rope_type", None) == "llama3":
|
||||
rope_scaling_factor = 1
|
||||
else:
|
||||
rope_scaling_factor = 1
|
||||
|
||||
for key in CONTEXT_LENGTH_KEYS:
|
||||
val = getattr(text_config, key, None)
|
||||
if val is not None:
|
||||
return int(rope_scaling_factor * val)
|
||||
return 2048
|
||||
|
||||
|
||||
# A fast LLaMA tokenizer with the pre-processed `tokenizer.json` file.
|
||||
_FAST_LLAMA_TOKENIZER = "hf-internal-testing/llama-tokenizer"
|
||||
|
||||
|
||||
def get_tokenizer(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast]:
|
||||
"""Gets a tokenizer for the given model name via Huggingface."""
|
||||
if tokenizer_name.endswith(".json"):
|
||||
from sglang.srt.tokenizer.tiktoken_tokenizer import TiktokenTokenizer
|
||||
|
||||
return TiktokenTokenizer(tokenizer_name)
|
||||
|
||||
if tokenizer_mode == "slow":
|
||||
if kwargs.get("use_fast", False):
|
||||
raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.")
|
||||
kwargs["use_fast"] = False
|
||||
|
||||
# TODO(Xinyuan): Remove this once we have a proper tokenizer for Devstral
|
||||
if tokenizer_name == "mistralai/Devstral-Small-2505":
|
||||
tokenizer_name = "mistralai/Mistral-Small-3.1-24B-Instruct-2503"
|
||||
|
||||
is_gguf = check_gguf_file(tokenizer_name)
|
||||
if is_gguf:
|
||||
kwargs["gguf_file"] = tokenizer_name
|
||||
tokenizer_name = Path(tokenizer_name).parent
|
||||
|
||||
if is_remote_url(tokenizer_name):
|
||||
# BaseConnector implements __del__() to clean up the local dir.
|
||||
# Since config files need to exist all the time, so we DO NOT use
|
||||
# with statement to avoid closing the client.
|
||||
client = create_remote_connector(tokenizer_name)
|
||||
client.pull_files(ignore_pattern=["*.pt", "*.safetensors", "*.bin"])
|
||||
tokenizer_name = client.get_local_dir()
|
||||
|
||||
try:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
tokenizer_revision=tokenizer_revision,
|
||||
clean_up_tokenization_spaces=False,
|
||||
**kwargs,
|
||||
)
|
||||
except TypeError as e:
|
||||
# The LLaMA tokenizer causes a protobuf error in some environments.
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If you are using a LLaMA V1 model "
|
||||
f"consider using '{_FAST_LLAMA_TOKENIZER}' instead of the "
|
||||
"original tokenizer."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
except ValueError as e:
|
||||
# If the error pertains to the tokenizer class not existing or not
|
||||
# currently being imported, suggest using the --trust-remote-code flag.
|
||||
if not trust_remote_code and (
|
||||
"does not exist or is not currently imported." in str(e)
|
||||
or "requires you to execute the tokenizer file" in str(e)
|
||||
):
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If the tokenizer is a custom "
|
||||
"tokenizer not yet available in the HuggingFace transformers "
|
||||
"library, consider setting `trust_remote_code=True` in LLM "
|
||||
"or using the `--trust-remote-code` flag in the CLI."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
else:
|
||||
raise e
|
||||
|
||||
if not isinstance(tokenizer, PreTrainedTokenizerFast):
|
||||
warnings.warn(
|
||||
"Using a slow tokenizer. This might cause a significant "
|
||||
"slowdown. Consider using a fast tokenizer instead."
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
return tokenizer
|
||||
|
||||
|
||||
# Some models doesn't have an available processor, e.g.: InternVL
|
||||
def get_tokenizer_from_processor(processor):
|
||||
if isinstance(processor, PreTrainedTokenizerBase):
|
||||
return processor
|
||||
return processor.tokenizer
|
||||
|
||||
|
||||
def get_processor(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
use_fast: Optional[bool] = True,
|
||||
**kwargs,
|
||||
):
|
||||
# pop 'revision' from kwargs if present.
|
||||
revision = kwargs.pop("revision", tokenizer_revision)
|
||||
|
||||
config = AutoConfig.from_pretrained(
|
||||
tokenizer_name,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# fix: for Qwen2-VL and Sarashina2Vision models, inject default 'size' if not provided.
|
||||
if config.model_type in {"qwen2_vl", "sarashina2_vision"}:
|
||||
if "size" not in kwargs:
|
||||
kwargs["size"] = {"shortest_edge": 3136, "longest_edge": 1003520}
|
||||
|
||||
if config.model_type not in {"llava", "clip"}:
|
||||
kwargs["use_fast"] = use_fast
|
||||
try:
|
||||
if "InternVL3_5" in tokenizer_name:
|
||||
processor = AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
except ValueError as e:
|
||||
error_message = str(e)
|
||||
if "does not have a slow version" in error_message:
|
||||
logger.info(
|
||||
f"Processor {tokenizer_name} does not have a slow version. Automatically use fast version"
|
||||
)
|
||||
kwargs["use_fast"] = True
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
raise e
|
||||
tokenizer = get_tokenizer_from_processor(processor)
|
||||
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
return processor
|
||||
|
||||
|
||||
def attach_additional_stop_token_ids(tokenizer):
|
||||
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
|
||||
if "<|eom_id|>" in tokenizer.get_added_vocab():
|
||||
tokenizer.additional_stop_token_ids = set(
|
||||
[tokenizer.get_added_vocab()["<|eom_id|>"]]
|
||||
)
|
||||
else:
|
||||
tokenizer.additional_stop_token_ids = None
|
||||
|
||||
|
||||
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
|
||||
"""Check if the file is a GGUF model."""
|
||||
model = Path(model)
|
||||
if not model.is_file():
|
||||
return False
|
||||
elif model.suffix == ".gguf":
|
||||
return True
|
||||
|
||||
with open(model, "rb") as f:
|
||||
header = f.read(4)
|
||||
return header == b"GGUF"
|
||||
@@ -0,0 +1,90 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
from typing import Callable, Union
|
||||
|
||||
import torch
|
||||
from packaging import version
|
||||
from torch.multiprocessing import reductions
|
||||
|
||||
from sglang.srt.utils import is_npu
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
|
||||
def monkey_patch_torch_reductions():
|
||||
"""Monkey patching before Torch https://github.com/pytorch/pytorch/pull/149248 is fixed"""
|
||||
|
||||
# Currently, NPU does not support UUID. This has been temporarily commented out, with support expected in the fourth quarter.
|
||||
if _is_npu:
|
||||
return
|
||||
|
||||
if hasattr(reductions, "_reduce_tensor_original"):
|
||||
return
|
||||
|
||||
reductions._reduce_tensor_original = reductions.reduce_tensor
|
||||
reductions._rebuild_cuda_tensor_original = reductions.rebuild_cuda_tensor
|
||||
|
||||
reductions.reduce_tensor = _reduce_tensor_modified
|
||||
reductions.rebuild_cuda_tensor = _rebuild_cuda_tensor_modified
|
||||
|
||||
reductions.init_reductions()
|
||||
|
||||
|
||||
# The signature has not been changed for years, and we will not need this when the next version is released,
|
||||
# so it looks safe to use a constant.
|
||||
_REDUCE_TENSOR_ARG_DEVICE_INDEX = 6
|
||||
|
||||
|
||||
def _reduce_tensor_modified(*args, **kwargs):
|
||||
output_fn, output_args = reductions._reduce_tensor_original(*args, **kwargs)
|
||||
output_args = _modify_tuple(
|
||||
output_args, _REDUCE_TENSOR_ARG_DEVICE_INDEX, _device_to_uuid
|
||||
)
|
||||
return output_fn, output_args
|
||||
|
||||
|
||||
def _rebuild_cuda_tensor_modified(*args):
|
||||
args = _modify_tuple(args, _REDUCE_TENSOR_ARG_DEVICE_INDEX, _device_from_maybe_uuid)
|
||||
return reductions._rebuild_cuda_tensor_original(*args)
|
||||
|
||||
|
||||
def _device_to_uuid(device: int) -> str:
|
||||
return str(torch.cuda.get_device_properties(device).uuid)
|
||||
|
||||
|
||||
def _device_from_maybe_uuid(device_maybe_uuid: Union[int, str]) -> int:
|
||||
if isinstance(device_maybe_uuid, int):
|
||||
return device_maybe_uuid
|
||||
|
||||
if isinstance(device_maybe_uuid, str):
|
||||
for device in range(torch.cuda.device_count()):
|
||||
if str(torch.cuda.get_device_properties(device).uuid) == device_maybe_uuid:
|
||||
return device
|
||||
raise Exception("Invalid device_uuid=" + device_maybe_uuid)
|
||||
|
||||
raise Exception(f"Unknown type: {device_maybe_uuid=}")
|
||||
|
||||
|
||||
def _modify_tuple(t, index: int, modifier: Callable):
|
||||
return *t[:index], modifier(t[index]), *t[index + 1 :]
|
||||
|
||||
|
||||
def monkey_patch_torch_compile():
|
||||
if version.parse(torch.__version__) < version.parse("2.8.0"):
|
||||
# These things are cacheable by torch.compile. torch.compile just doesn't know it.
|
||||
# This was fixed in PyTorch 2.8, but until then, we monkey patch.
|
||||
import torch._higher_order_ops.auto_functionalize as af
|
||||
|
||||
af.auto_functionalized_v2._cacheable = True
|
||||
af.auto_functionalized._cacheable = True
|
||||
@@ -0,0 +1,31 @@
|
||||
import torch
|
||||
|
||||
from sglang.srt.distributed import get_world_group
|
||||
|
||||
|
||||
class PollBasedBarrier:
|
||||
def __init__(self, noop: bool = False):
|
||||
self._noop = noop
|
||||
self._local_arrived = False
|
||||
|
||||
def local_arrive(self):
|
||||
assert not self._local_arrived
|
||||
self._local_arrived = True
|
||||
|
||||
def poll_global_arrived(self) -> bool:
|
||||
global_arrived = self._compute_global_arrived()
|
||||
output = self._local_arrived and global_arrived
|
||||
if output:
|
||||
self._local_arrived = False
|
||||
return output
|
||||
|
||||
def _compute_global_arrived(self) -> bool:
|
||||
local_arrived = self._noop or self._local_arrived
|
||||
global_arrived = torch.tensor(local_arrived)
|
||||
# Can optimize if bottleneck
|
||||
torch.distributed.all_reduce(
|
||||
global_arrived,
|
||||
torch.distributed.ReduceOp.MIN,
|
||||
group=get_world_group().cpu_group,
|
||||
)
|
||||
return global_arrived.item()
|
||||
@@ -0,0 +1,452 @@
|
||||
# https://raw.githubusercontent.com/ROCm/rocmProfileData/refs/heads/master/tools/rpd2tracing.py
|
||||
# commit 92d13a08328625463e9ba944cece82fc5eea36e6
|
||||
def rpd_to_chrome_trace(
|
||||
input_rpd, output_json=None, start="0%", end="100%", format="object"
|
||||
):
|
||||
import gzip
|
||||
import sqlite3
|
||||
|
||||
if output_json is None:
|
||||
import pathlib
|
||||
|
||||
output_json = pathlib.PurePath(input_rpd).with_suffix(".trace.json.gz")
|
||||
|
||||
connection = sqlite3.connect(input_rpd)
|
||||
|
||||
outfile = gzip.open(output_json, "wt", encoding="utf-8")
|
||||
|
||||
if format == "object":
|
||||
outfile.write('{"traceEvents": ')
|
||||
|
||||
outfile.write("[ {}\n")
|
||||
|
||||
for row in connection.execute("select distinct gpuId from rocpd_op"):
|
||||
try:
|
||||
outfile.write(
|
||||
',{"name": "process_name", "ph": "M", "pid":"%s","args":{"name":"%s"}}\n'
|
||||
% (row[0], "GPU" + str(row[0]))
|
||||
)
|
||||
outfile.write(
|
||||
',{"name": "process_sort_index", "ph": "M", "pid":"%s","args":{"sort_index":"%s"}}\n'
|
||||
% (row[0], row[0] + 1000000)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
for row in connection.execute("select distinct pid, tid from rocpd_api"):
|
||||
try:
|
||||
outfile.write(
|
||||
',{"name":"thread_name","ph":"M","pid":"%s","tid":"%s","args":{"name":"%s"}}\n'
|
||||
% (row[0], row[1], "Hip " + str(row[1]))
|
||||
)
|
||||
outfile.write(
|
||||
',{"name":"thread_sort_index","ph":"M","pid":"%s","tid":"%s","args":{"sort_index":"%s"}}\n'
|
||||
% (row[0], row[1], row[1] * 2)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
try:
|
||||
# FIXME - these aren't rendering correctly in chrome://tracing
|
||||
for row in connection.execute("select distinct pid, tid from rocpd_hsaApi"):
|
||||
try:
|
||||
outfile.write(
|
||||
',{"name":"thread_name","ph":"M","pid":"%s","tid":"%s","args":{"name":"%s"}}\n'
|
||||
% (row[0], row[1], "HSA " + str(row[1]))
|
||||
)
|
||||
outfile.write(
|
||||
',{"name":"thread_sort_index","ph":"M","pid":"%s","tid":"%s","args":{"sort_index":"%s"}}\n'
|
||||
% (row[0], row[1], row[1] * 2 - 1)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
except:
|
||||
pass
|
||||
|
||||
rangeStringApi = ""
|
||||
rangeStringOp = ""
|
||||
rangeStringMonitor = ""
|
||||
min_time = connection.execute("select MIN(start) from rocpd_api;").fetchall()[0][0]
|
||||
max_time = connection.execute("select MAX(end) from rocpd_api;").fetchall()[0][0]
|
||||
if min_time == None:
|
||||
raise Exception("Trace file is empty.")
|
||||
|
||||
print("Timestamps:")
|
||||
print(f"\t first: \t{min_time/1000} us")
|
||||
print(f"\t last: \t{max_time/1000} us")
|
||||
print(f"\t duration: \t{(max_time-min_time) / 1000000000} seconds")
|
||||
|
||||
start_time = min_time / 1000
|
||||
end_time = max_time / 1000
|
||||
|
||||
if start:
|
||||
if "%" in start:
|
||||
start_time = (
|
||||
(max_time - min_time) * (int(start.replace("%", "")) / 100) + min_time
|
||||
) / 1000
|
||||
else:
|
||||
start_time = int(start)
|
||||
rangeStringApi = "where rocpd_api.start/1000 >= %s" % (start_time)
|
||||
rangeStringOp = "where rocpd_op.start/1000 >= %s" % (start_time)
|
||||
rangeStringMonitor = "where start/1000 >= %s" % (start_time)
|
||||
if end:
|
||||
if "%" in end:
|
||||
end_time = (
|
||||
(max_time - min_time) * (int(end.replace("%", "")) / 100) + min_time
|
||||
) / 1000
|
||||
else:
|
||||
end_time = int(end)
|
||||
|
||||
rangeStringApi = (
|
||||
rangeStringApi + " and rocpd_api.start/1000 <= %s" % (end_time)
|
||||
if start != None
|
||||
else "where rocpd_api.start/1000 <= %s" % (end_time)
|
||||
)
|
||||
rangeStringOp = (
|
||||
rangeStringOp + " and rocpd_op.start/1000 <= %s" % (end_time)
|
||||
if start != None
|
||||
else "where rocpd_op.start/1000 <= %s" % (end_time)
|
||||
)
|
||||
rangeStringMonitor = (
|
||||
rangeStringMonitor + " and start/1000 <= %s" % (end_time)
|
||||
if start != None
|
||||
else "where start/1000 <= %s" % (end_time)
|
||||
)
|
||||
|
||||
print("\nFilter: %s" % (rangeStringApi))
|
||||
print(f"Output duration: {(end_time-start_time)/1000000} seconds")
|
||||
|
||||
# Output Ops
|
||||
|
||||
for row in connection.execute(
|
||||
"select A.string as optype, B.string as description, gpuId, queueId, rocpd_op.start/1000.0, (rocpd_op.end-rocpd_op.start) / 1000.0 from rocpd_op INNER JOIN rocpd_string A on A.id = rocpd_op.opType_id INNER Join rocpd_string B on B.id = rocpd_op.description_id %s"
|
||||
% (rangeStringOp)
|
||||
):
|
||||
try:
|
||||
name = row[0] if len(row[1]) == 0 else row[1]
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'
|
||||
% (row[2], row[3], name, row[4], row[5], row[0])
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
# Output Graph executions on GPU
|
||||
try:
|
||||
for row in connection.execute(
|
||||
"select graphExec, gpuId, queueId, min(start)/1000.0, (max(end)-min(start))/1000.0, count(*) from rocpd_graphLaunchapi A join rocpd_api_ops B on B.api_id = A.api_ptr_id join rocpd_op C on C.id = B.op_id %s group by api_ptr_id"
|
||||
% (rangeStringMonitor)
|
||||
):
|
||||
try:
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"kernels":"%s"}}\n'
|
||||
% (row[1], row[2], f"Graph {row[0]}", row[3], row[4], row[5])
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
except:
|
||||
pass
|
||||
|
||||
# Output apis
|
||||
for row in connection.execute(
|
||||
"select A.string as apiName, B.string as args, pid, tid, rocpd_api.start/1000.0, (rocpd_api.end-rocpd_api.start) / 1000.0, (rocpd_api.end != rocpd_api.start) as has_duration from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id INNER Join rocpd_string B on B.id = rocpd_api.args_id %s order by rocpd_api.id"
|
||||
% (rangeStringApi)
|
||||
):
|
||||
try:
|
||||
if row[0] == "UserMarker":
|
||||
if row[6] == 0: # instantanuous "mark" messages
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","ph":"i","s":"p","args":{"desc":"%s"}}\n'
|
||||
% (
|
||||
row[2],
|
||||
row[3],
|
||||
row[1].replace('"', ""),
|
||||
row[4],
|
||||
row[1].replace('"', ""),
|
||||
)
|
||||
)
|
||||
else:
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'
|
||||
% (
|
||||
row[2],
|
||||
row[3],
|
||||
row[1].replace('"', ""),
|
||||
row[4],
|
||||
row[5],
|
||||
row[1].replace('"', ""),
|
||||
)
|
||||
)
|
||||
else:
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'
|
||||
% (
|
||||
row[2],
|
||||
row[3],
|
||||
row[0],
|
||||
row[4],
|
||||
row[5],
|
||||
row[1].replace('"', "").replace("\t", ""),
|
||||
)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
# Output api->op linkage
|
||||
for row in connection.execute(
|
||||
"select rocpd_api_ops.id, pid, tid, gpuId, queueId, rocpd_api.end/1000.0 - 2, rocpd_op.start/1000.0 from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id %s"
|
||||
% (rangeStringApi)
|
||||
):
|
||||
try:
|
||||
fromtime = row[5] if row[5] < row[6] else row[6]
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","cat":"api_op","name":"api_op","ts":"%s","id":"%s","ph":"s"}\n'
|
||||
% (row[1], row[2], fromtime, row[0])
|
||||
)
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","cat":"api_op","name":"api_op","ts":"%s","id":"%s","ph":"f", "bp":"e"}\n'
|
||||
% (row[3], row[4], row[6], row[0])
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
try:
|
||||
for row in connection.execute(
|
||||
"select A.string as apiName, B.string as args, pid, tid, rocpd_hsaApi.start/1000.0, (rocpd_hsaApi.end-rocpd_hsaApi.start) / 1000.0 from rocpd_hsaApi INNER JOIN rocpd_string A on A.id = rocpd_hsaApi.apiName_id INNER Join rocpd_string B on B.id = rocpd_hsaApi.args_id %s order by rocpd_hsaApi.id"
|
||||
% (rangeStringApi)
|
||||
):
|
||||
try:
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'
|
||||
% (
|
||||
row[2],
|
||||
row[3] + 1,
|
||||
row[0],
|
||||
row[4],
|
||||
row[5],
|
||||
row[1].replace('"', ""),
|
||||
)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
except:
|
||||
pass
|
||||
|
||||
#
|
||||
# Counters
|
||||
#
|
||||
|
||||
# Counters should extend to the last event in the trace. This means they need to have a value at Tend.
|
||||
# Figure out when that is
|
||||
|
||||
T_end = 0
|
||||
for row in connection.execute(
|
||||
"SELECT max(end)/1000 from (SELECT end from rocpd_api UNION ALL SELECT end from rocpd_op)"
|
||||
):
|
||||
T_end = int(row[0])
|
||||
if end:
|
||||
T_end = end_time
|
||||
|
||||
# Loop over GPU for per-gpu counters
|
||||
gpuIdsPresent = []
|
||||
for row in connection.execute("SELECT DISTINCT gpuId FROM rocpd_op"):
|
||||
gpuIdsPresent.append(row[0])
|
||||
|
||||
for gpuId in gpuIdsPresent:
|
||||
# print(f"Creating counters for: {gpuId}")
|
||||
|
||||
# Create the queue depth counter
|
||||
depth = 0
|
||||
idle = 1
|
||||
for row in connection.execute(
|
||||
'select * from (select rocpd_api.start/1000.0 as ts, "1" from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id AND rocpd_op.gpuId = %s %s UNION ALL select rocpd_op.end/1000.0, "-1" from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id AND rocpd_op.gpuId = %s %s) order by ts'
|
||||
% (gpuId, rangeStringOp, gpuId, rangeStringOp)
|
||||
):
|
||||
try:
|
||||
if idle and int(row[1]) > 0:
|
||||
idle = 0
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'
|
||||
% (gpuId, row[0], idle)
|
||||
)
|
||||
if depth == 1 and int(row[1]) < 0:
|
||||
idle = 1
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'
|
||||
% (gpuId, row[0], idle)
|
||||
)
|
||||
depth = depth + int(row[1])
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"QueueDepth","ph":"C","ts":%s,"args":{"depth":%s}}\n'
|
||||
% (gpuId, row[0], depth)
|
||||
)
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
if T_end > 0:
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'
|
||||
% (gpuId, T_end, idle)
|
||||
)
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"QueueDepth","ph":"C","ts":%s,"args":{"depth":%s}}\n'
|
||||
% (gpuId, T_end, depth)
|
||||
)
|
||||
|
||||
# Create SMI counters
|
||||
try:
|
||||
for row in connection.execute(
|
||||
"select deviceId, monitorType, start/1000.0, value from rocpd_monitor %s"
|
||||
% (rangeStringMonitor)
|
||||
):
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"%s","ph":"C","ts":%s,"args":{"%s":%s}}\n'
|
||||
% (row[0], row[1], row[2], row[1], row[3])
|
||||
)
|
||||
# Output the endpoints of the last range
|
||||
for row in connection.execute(
|
||||
"select distinct deviceId, monitorType, max(end)/1000.0, value from rocpd_monitor %s group by deviceId, monitorType"
|
||||
% (rangeStringMonitor)
|
||||
):
|
||||
outfile.write(
|
||||
',{"pid":"%s","name":"%s","ph":"C","ts":%s,"args":{"%s":%s}}\n'
|
||||
% (row[0], row[1], row[2], row[1], row[3])
|
||||
)
|
||||
except:
|
||||
print("Did not find SMI data")
|
||||
|
||||
# Create the (global) memory counter
|
||||
"""
|
||||
sizes = {} # address -> size
|
||||
totalSize = 0
|
||||
exp = re.compile("^ptr\((.*)\)\s+size\((.*)\)$")
|
||||
exp2 = re.compile("^ptr\((.*)\)$")
|
||||
for row in connection.execute("SELECT rocpd_api.end/1000.0 as ts, B.string, '1' FROM rocpd_api INNER JOIN rocpd_string A ON A.id=rocpd_api.apiName_id INNER JOIN rocpd_string B ON B.id=rocpd_api.args_id WHERE A.string='hipFree' UNION ALL SELECT rocpd_api.start/1000.0, B.string, '0' FROM rocpd_api INNER JOIN rocpd_string A ON A.id=rocpd_api.apiName_id INNER JOIN rocpd_string B ON B.id=rocpd_api.args_id WHERE A.string='hipMalloc' ORDER BY ts asc"):
|
||||
try:
|
||||
if row[2] == '0': #malloc
|
||||
m = exp.match(row[1])
|
||||
if m:
|
||||
size = int(m.group(2), 16)
|
||||
totalSize = totalSize + size
|
||||
sizes[m.group(1)] = size
|
||||
outfile.write(',{"pid":"0","name":"Allocated Memory","ph":"C","ts":%s,"args":{"depth":%s}}\n'%(row[0],totalSize))
|
||||
else: #free
|
||||
m = exp2.match(row[1])
|
||||
if m:
|
||||
try: # Sometimes free addresses are not valid or listed
|
||||
size = sizes[m.group(1)]
|
||||
sizes[m.group(1)] = 0
|
||||
totalSize = totalSize - size;
|
||||
outfile.write(',{"pid":"0","name":"Allocated Memory","ph":"C","ts":%s,"args":{"depth":%s}}\n'%(row[0],totalSize))
|
||||
except KeyError:
|
||||
pass
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
if T_end > 0:
|
||||
outfile.write(',{"pid":"0","name":"Allocated Memory","ph":"C","ts":%s,"args":{"depth":%s}}\n'%(T_end,totalSize))
|
||||
"""
|
||||
|
||||
# Create "faux calling stack frame" on gpu ops traceS
|
||||
stacks = {} # Call stacks built from UserMarker entres. Key is 'pid,tid'
|
||||
currentFrame = {} # "Current GPU frame" (id, name, start, end). Key is 'pid,tid'
|
||||
|
||||
class GpuFrame:
|
||||
def __init__(self):
|
||||
self.id = 0
|
||||
self.name = ""
|
||||
self.start = 0
|
||||
self.end = 0
|
||||
self.gpus = []
|
||||
self.totalOps = 0
|
||||
|
||||
# FIXME: include 'start' (in ns) so we can ORDER BY it and break ties?
|
||||
for row in connection.execute(
|
||||
"SELECT '0', start/1000.0, pid, tid, B.string as label, '','','', '' from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id AND A.string = 'UserMarker' INNER JOIN rocpd_string B on B.id = rocpd_api.args_id AND rocpd_api.start/1000.0 != rocpd_api.end/1000.0 %s UNION ALL SELECT '1', end/1000.0, pid, tid, B.string as label, '','','', '' from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id AND A.string = 'UserMarker' INNER JOIN rocpd_string B on B.id = rocpd_api.args_id AND rocpd_api.start/1000.0 != rocpd_api.end/1000.0 %s UNION ALL SELECT '2', rocpd_api.start/1000.0, pid, tid, '' as label, gpuId, queueId, rocpd_op.start/1000.0, rocpd_op.end/1000.0 from rocpd_api_ops INNER JOIN rocpd_api ON rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op ON rocpd_api_ops.op_id = rocpd_op.id %s ORDER BY start/1000.0 asc"
|
||||
% (rangeStringApi, rangeStringApi, rangeStringApi)
|
||||
):
|
||||
try:
|
||||
key = (row[2], row[3]) # Key is 'pid,tid'
|
||||
if row[0] == "0": # Frame start
|
||||
if key not in stacks:
|
||||
stacks[key] = []
|
||||
stack = stacks[key].append((row[1], row[4]))
|
||||
# print(f"0: new api frame: pid_tid={key} -> stack={stacks}")
|
||||
|
||||
elif row[0] == "1": # Frame end
|
||||
completed = stacks[key].pop()
|
||||
# print(f"1: end api frame: pid_tid={key} -> stack={stacks}")
|
||||
|
||||
elif row[0] == "2": # API + Op
|
||||
if key in stacks and len(stacks[key]) > 0:
|
||||
frame = stacks[key][-1]
|
||||
# print(f"2: Op on {frame} ({len(stacks[key])})")
|
||||
gpuFrame = None
|
||||
if key not in currentFrame: # First op under the current api frame
|
||||
gpuFrame = GpuFrame()
|
||||
gpuFrame.id = frame[0]
|
||||
gpuFrame.name = frame[1]
|
||||
gpuFrame.start = row[7]
|
||||
gpuFrame.end = row[8]
|
||||
gpuFrame.gpus.append((row[5], row[6]))
|
||||
gpuFrame.totalOps = 1
|
||||
# print(f"2a: new frame: {gpuFrame.gpus} {gpuFrame.start} {gpuFrame.end} {gpuFrame.end - gpuFrame.start}")
|
||||
else:
|
||||
gpuFrame = currentFrame[key]
|
||||
# Another op under the same frame -> union them (but only if they are butt together)
|
||||
if (
|
||||
gpuFrame.id == frame[0]
|
||||
and gpuFrame.name == frame[1]
|
||||
and (
|
||||
abs(row[7] - gpuFrame.end) < 200
|
||||
or abs(gpuFrame.start - row[8]) < 200
|
||||
)
|
||||
):
|
||||
# if gpuFrame.id == frame[0] and gpuFrame.name == frame[1]: # Another op under the same frame -> union them
|
||||
# if False: # Turn off frame joining
|
||||
if row[7] < gpuFrame.start:
|
||||
gpuFrame.start = row[7]
|
||||
if row[8] > gpuFrame.end:
|
||||
gpuFrame.end = row[8]
|
||||
if (row[5], row[6]) not in gpuFrame.gpus:
|
||||
gpuFrame.gpus.append((row[5], row[6]))
|
||||
gpuFrame.totalOps = gpuFrame.totalOps + 1
|
||||
# print(f"2c: union frame: {gpuFrame.gpus} {gpuFrame.start} {gpuFrame.end} {gpuFrame.end - gpuFrame.start}")
|
||||
|
||||
else: # This is a new frame - dump the last and make new
|
||||
gpuFrame = currentFrame[key]
|
||||
for dest in gpuFrame.gpus:
|
||||
# print(f"2: OUTPUT: dest={dest} time={gpuFrame.start} -> {gpuFrame.end} Duration={gpuFrame.end - gpuFrame.start} TotalOps={gpuFrame.totalOps}")
|
||||
outfile.write(
|
||||
',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'
|
||||
% (
|
||||
dest[0],
|
||||
dest[1],
|
||||
gpuFrame.name.replace('"', ""),
|
||||
gpuFrame.start - 1,
|
||||
gpuFrame.end - gpuFrame.start + 1,
|
||||
f"UserMarker frame: {gpuFrame.totalOps} ops",
|
||||
)
|
||||
)
|
||||
currentFrame.pop(key)
|
||||
|
||||
# make the first op under the new frame
|
||||
gpuFrame = GpuFrame()
|
||||
gpuFrame.id = frame[0]
|
||||
gpuFrame.name = frame[1]
|
||||
gpuFrame.start = row[7]
|
||||
gpuFrame.end = row[8]
|
||||
gpuFrame.gpus.append((row[5], row[6]))
|
||||
gpuFrame.totalOps = 1
|
||||
# print(f"2b: new frame: {gpuFrame.gpus} {gpuFrame.start} {gpuFrame.end} {gpuFrame.end - gpuFrame.start}")
|
||||
|
||||
currentFrame[key] = gpuFrame
|
||||
|
||||
except ValueError:
|
||||
outfile.write("")
|
||||
|
||||
outfile.write("]\n")
|
||||
|
||||
if format == "object":
|
||||
outfile.write("} \n")
|
||||
|
||||
outfile.close()
|
||||
connection.close()
|
||||
@@ -0,0 +1,71 @@
|
||||
import logging
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import triton
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def execute():
|
||||
if dist.get_rank() == 0:
|
||||
logger.info(f"[slow_rank_detector] Start benchmarking...")
|
||||
|
||||
local_metrics = {
|
||||
bench_name: _compute_local_metric(bench_name) for bench_name in _BENCH_NAMES
|
||||
}
|
||||
|
||||
all_metrics = [None for _ in range(dist.get_world_size())]
|
||||
dist.gather_object(local_metrics, all_metrics if dist.get_rank() == 0 else None)
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
_analyze_metrics(all_metrics)
|
||||
|
||||
|
||||
class _GemmExecutor:
|
||||
def __init__(self):
|
||||
self.lhs = torch.randn((8192, 8192), dtype=torch.bfloat16, device="cuda")
|
||||
self.rhs = torch.randn((8192, 8192), dtype=torch.bfloat16, device="cuda")
|
||||
|
||||
def __call__(self):
|
||||
self.lhs @ self.rhs
|
||||
|
||||
|
||||
class _ElementwiseExecutor:
|
||||
def __init__(self):
|
||||
self.value = torch.randint(
|
||||
0, 10000, (128 * 1024**2,), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
def __call__(self):
|
||||
self.value += 1
|
||||
|
||||
|
||||
_EXECUTOR_CLS_OF_BENCH = {
|
||||
"gemm": _GemmExecutor,
|
||||
"elementwise": _ElementwiseExecutor,
|
||||
}
|
||||
|
||||
_BENCH_NAMES = list(_EXECUTOR_CLS_OF_BENCH.keys())
|
||||
|
||||
|
||||
def _compute_local_metric(bench_name):
|
||||
executor = _EXECUTOR_CLS_OF_BENCH[bench_name]()
|
||||
ms = triton.testing.do_bench_cudagraph(executor, return_mode="mean", rep=20)
|
||||
return ms
|
||||
|
||||
|
||||
def _analyze_metrics(all_metrics: List[Dict[str, Any]]):
|
||||
for bench_name in _BENCH_NAMES:
|
||||
time_of_rank = torch.tensor([m[bench_name] for m in all_metrics])
|
||||
speed_of_rank = 1 / time_of_rank
|
||||
rel_speed_of_rank = speed_of_rank / speed_of_rank.max()
|
||||
slowest_rel_speed = rel_speed_of_rank.min().item()
|
||||
logger.info(
|
||||
f"[slow_rank_detector] {bench_name=} {slowest_rel_speed=} {rel_speed_of_rank=} {time_of_rank=}"
|
||||
)
|
||||
if slowest_rel_speed < 0.9:
|
||||
logger.warning(
|
||||
"[slow_rank_detector] Some ranks are too slow compared with others"
|
||||
)
|
||||
Reference in New Issue
Block a user