mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
remove slime/tests folder
This commit is contained in:
@@ -1,120 +0,0 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
import pytest
|
||||
import ray
|
||||
|
||||
CKPT_ARGS=[
|
||||
"--hf-checkpoint", "/root/Qwen3-0.6B",
|
||||
]
|
||||
|
||||
|
||||
ROLLOUT_ARGS=[
|
||||
"--prompt-data", "/root/dapo-math-17k/dapo-math-17k.jsonl",
|
||||
"--input-key", "prompt",
|
||||
"--label-key", "label",
|
||||
"--apply-chat-template",
|
||||
"--rollout-shuffle",
|
||||
"--rm-type", "deepscaler",
|
||||
"--num-rollout", "3000",
|
||||
"--rollout-batch-size", "16",
|
||||
"--n-samples-per-prompt", "16",
|
||||
"--rollout-max-response-len", "8192",
|
||||
"--rollout-temperature", "0.8",
|
||||
"--global-batch-size", "128",
|
||||
]
|
||||
|
||||
GRPO_ARGS=[
|
||||
"--advantage-estimator", "grpo",
|
||||
"--kl-loss-coef", "0.00",
|
||||
"--kl-loss-type", "low_var_kl",
|
||||
"--kl-coef", "0.00",
|
||||
"--entropy-coef", "0.00",
|
||||
"--eps-clip", "0.2",
|
||||
"--eps-clip-high", "0.28",
|
||||
]
|
||||
|
||||
OPTIMIZER_ARGS=[
|
||||
"--optimizer", "adam",
|
||||
"--lr", "1e-6",
|
||||
"--lr-decay-style", "constant",
|
||||
"--weight-decay", "0.1",
|
||||
"--adam-beta1", "0.9",
|
||||
"--adam-beta2", "0.98",
|
||||
]
|
||||
|
||||
SGLANG_ARGS=[
|
||||
"--rollout-num-gpus-per-engine", "1",
|
||||
"--load-format", "dummy",
|
||||
]
|
||||
|
||||
TRAIN_ARGS=[
|
||||
"--actor-num-nodes", "1",
|
||||
"--actor-num-gpus-per-node", "4",
|
||||
]
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
def setup_env():
|
||||
os.environ["PYTHONBUFFERED"] = "16"
|
||||
# TODO: Temporary workaround to assign environment variables both outside and inside Ray.
|
||||
os.environ["no_proxy"] = "localhost,127.0.0.1,0.0.0.0"
|
||||
os.environ["SLIME_BACKEND"] = "fsdp"
|
||||
|
||||
# Ref: https://github.com/ray-project/ray/blob/ray-2.49.1/python/ray/tests/conftest.py
|
||||
def get_default_fixure_system_config():
|
||||
system_config = {
|
||||
"object_timeout_milliseconds": 200,
|
||||
"health_check_initial_delay_ms": 0,
|
||||
"health_check_failure_threshold": 10,
|
||||
"object_store_full_delay_ms": 100,
|
||||
"local_gc_min_interval_s": 1,
|
||||
}
|
||||
return system_config
|
||||
|
||||
def get_default_fixture_ray_kwargs():
|
||||
system_config = get_default_fixure_system_config()
|
||||
ray_kwargs = {
|
||||
"num_cpus": 1,
|
||||
"object_store_memory": 150 * 1024 * 1024,
|
||||
"dashboard_port": None,
|
||||
"namespace": "default_test_namespace",
|
||||
"_system_config": system_config,
|
||||
}
|
||||
return ray_kwargs
|
||||
|
||||
def get_slime_ray_kwargs():
|
||||
ray_kwargs = {
|
||||
"runtime_env": {
|
||||
"env_vars": {
|
||||
"no_proxy": "localhost,127.0.0.1,0.0.0.0",
|
||||
"SLIME_BACKEND": "fsdp",
|
||||
}
|
||||
}
|
||||
}
|
||||
return ray_kwargs
|
||||
|
||||
@contextmanager
|
||||
def _ray_start(**kwargs):
|
||||
init_kwargs = get_default_fixture_ray_kwargs()
|
||||
init_kwargs.update(kwargs)
|
||||
init_kwargs.update(get_slime_ray_kwargs())
|
||||
# Start the Ray processes.
|
||||
address_info = ray.init("local", **init_kwargs)
|
||||
|
||||
yield address_info
|
||||
# The code after the yield will run as teardown code.
|
||||
ray.shutdown()
|
||||
# Delete the cluster address just in case.
|
||||
ray._common.utils.reset_ray_address()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ray_start_regular(request):
|
||||
param = getattr(request, "param", {})
|
||||
with _ray_start(**param) as res:
|
||||
yield res
|
||||
|
||||
@pytest.fixture
|
||||
def ray_start_4_gpus_unlimited_cpus(request):
|
||||
param = getattr(request, "param", {})
|
||||
with _ray_start(num_cpus=None, num_gpus=4, object_store_memory=None, **param) as res:
|
||||
yield res
|
||||
File diff suppressed because one or more lines are too long
@@ -1,128 +0,0 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import transformer_engine_torch as tex
|
||||
import triton
|
||||
from transformer_engine.pytorch.tensor.float8_blockwise_tensor import Float8BlockQuantizer
|
||||
|
||||
from slime.utils.fp8_kernel import blockwise_cast_to_fp8_triton
|
||||
|
||||
device = "cuda"
|
||||
dtype = torch.bfloat16
|
||||
fp8_dtype = torch.float8_e4m3fn
|
||||
fp8_max = torch.finfo(fp8_dtype).max
|
||||
fp8_min = -fp8_max
|
||||
|
||||
|
||||
def ceil_div(x: int, y: int) -> int:
|
||||
"""
|
||||
Perform ceiling division of two integers.
|
||||
|
||||
Args:
|
||||
x: the dividend.
|
||||
y: the divisor.
|
||||
|
||||
Returns:
|
||||
The result of the ceiling division.
|
||||
"""
|
||||
return (x + y - 1) // y
|
||||
|
||||
|
||||
def per_block_cast_to_fp8_slime(weight, weight_block_size=[128, 128]):
|
||||
FP8_MIN = torch.finfo(torch.float8_e4m3fn).min
|
||||
FP8_MAX = torch.finfo(torch.float8_e4m3fn).max
|
||||
|
||||
# per block quant
|
||||
block_n, block_k = weight_block_size[0], weight_block_size[1]
|
||||
|
||||
shape_0, shape_1 = weight.shape
|
||||
|
||||
n_tiles = ceil_div(shape_0, block_n)
|
||||
k_tiles = ceil_div(shape_1, block_k)
|
||||
|
||||
q_weight = F.pad(
|
||||
weight,
|
||||
(0, k_tiles * block_k - shape_1, 0, n_tiles * block_n - shape_0),
|
||||
mode="constant",
|
||||
value=0.0,
|
||||
)
|
||||
|
||||
qweight = q_weight.reshape(n_tiles, block_n, k_tiles, block_k)
|
||||
block_max = torch.max(torch.abs(qweight), dim=1, keepdim=True)[0]
|
||||
block_max = torch.max(block_max, dim=3, keepdim=True)[0]
|
||||
|
||||
scale = block_max.to(torch.float32) / FP8_MAX
|
||||
qweight = (
|
||||
(qweight / scale)
|
||||
.clamp(min=FP8_MIN, max=FP8_MAX)
|
||||
.reshape((n_tiles * block_n, k_tiles * block_k))
|
||||
.to(torch.float8_e4m3fn)
|
||||
)
|
||||
qweight = qweight[:shape_0, :shape_1]
|
||||
scale = scale.squeeze()
|
||||
return qweight, scale
|
||||
|
||||
|
||||
def te_per_token_group_quant_8bit(weight: torch.Tensor, quantizer, weight_block_size=[128, 128]):
|
||||
block_n, block_k = weight_block_size[0], weight_block_size[1]
|
||||
shape_0, shape_1 = weight.shape
|
||||
n_tiles = ceil_div(shape_0, block_n)
|
||||
k_tiles = ceil_div(shape_1, block_k)
|
||||
param = quantizer(weight)
|
||||
return param._rowwise_data, param._rowwise_scale_inv[:n_tiles, :k_tiles]
|
||||
|
||||
|
||||
ref_lib = "pytorch" # pytorch or te
|
||||
configs = []
|
||||
configs.append(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["M", "N"], # Argument names to use as an x-axis for the plot
|
||||
x_vals=[128 * i for i in range(2, 33)], # Different possible values for `x_name`
|
||||
line_arg="provider", # Argument name whose value corresponds to a different line in the plot
|
||||
# Possible values for `line_arg`
|
||||
# Don't compare to cublas for fp8 cases as torch.matmul doesn't support fp8 at the moment.
|
||||
line_vals=[ref_lib, "triton"], # Label name for the lines
|
||||
line_names=[ref_lib, "Triton"], # Line styles
|
||||
styles=[("green", "-"), ("blue", "-")],
|
||||
ylabel="GB/s", # Label name for the y-axis
|
||||
plot_name="quant-performance", # Name for the plot, used also as a file name for saving the plot.
|
||||
args={}, # Values for function arguments not in `x_names` and `y_name`.
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(configs)
|
||||
def benchmark(M, N, provider):
|
||||
x = torch.randn((M, N), device=device, dtype=dtype)
|
||||
quantiles = [0.5, 0.2, 0.8]
|
||||
if provider == "pytorch":
|
||||
ms, min_ms, max_ms = triton.testing.do_bench(lambda: per_block_cast_to_fp8_slime(x), quantiles=quantiles)
|
||||
if provider == "triton":
|
||||
ms, min_ms, max_ms = triton.testing.do_bench(lambda: blockwise_cast_to_fp8_triton(x), quantiles=quantiles)
|
||||
if provider == "te":
|
||||
quantizer = Float8BlockQuantizer(
|
||||
fp8_dtype=tex.DType.kFloat8E4M3,
|
||||
rowwise=True,
|
||||
columnwise=True,
|
||||
force_pow_2_scales=False,
|
||||
block_scaling_dim=2,
|
||||
)
|
||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
||||
lambda: te_per_token_group_quant_8bit(x, quantizer), quantiles=quantiles
|
||||
)
|
||||
gbps = lambda ms: x.numel() * x.element_size() * 1e-9 / (ms * 1e-3)
|
||||
return gbps(ms), gbps(max_ms), gbps(min_ms)
|
||||
|
||||
|
||||
def benchmark_percise():
|
||||
for M in (7168, 2112, 1536, 24576, 512, 32768, 16384, 4096, 2048):
|
||||
for N in (2048, 4096, 8192):
|
||||
x_ref = torch.rand(M, N, dtype=dtype, device=device)
|
||||
x_triton, x_s_triton = blockwise_cast_to_fp8_triton(x_ref)
|
||||
x_slime, x_s_slime = per_block_cast_to_fp8_slime(x_ref)
|
||||
torch.testing.assert_close(x_triton.to(torch.float32), x_slime.to(torch.float32), rtol=1e-3, atol=1e-5)
|
||||
torch.testing.assert_close(x_s_triton, x_s_slime, rtol=1e-3, atol=1e-5)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark_percise()
|
||||
benchmark.run(show_plots=True, print_data=True, save_path=f"./plot/{ref_lib}")
|
||||
Reference in New Issue
Block a user