remove slime/tests folder

This commit is contained in:
Zilin Zhu
2025-09-27 01:34:53 +00:00
parent 4dd5e7dfb5
commit d862870c6a
4 changed files with 0 additions and 319 deletions
-120
View File
@@ -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
-128
View File
@@ -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}")