Skip to main content

add-reference-tests

Add pytest tests to validate reference implementations in the flashinfer-trace HuggingFace dataset against FlashInfer or SGLang ground truth. Use when validating kernel definitions, adding tests for new op_types, or verifying reference implementations are correct.

الانتقال إلى التثبيت

معلومات المصدر

المستودع
flashinfer-ai/flashinfer-bench
آخر نشاط في المصدر
٢٧ أبريل ٢٠٢٦ في ٠٦:٢٥
لغة SKILL.md المكتشفة
الإنجليزية
النجوم
٢٨١
التفرعات
٥٠

خيارات التثبيت

يُحدَّد Prompt الذي يراجع المصدر أولًا بشكل افتراضي. يمكنك التبديل إلى أمر مباشر أو تنزيل نسخة محلية.

مراجعة ملفات المصدر

اقرأ SKILL.md وأي ملفات مرافقة يعرضها SkillsMP قبل أن تقرر التثبيت.

عرض SKILL.md

SKILL.md
تعليمات المصدر · معاينة للقراءة فقط
name
add-reference-tests
description
Add pytest tests to validate reference implementations in the flashinfer-trace HuggingFace dataset against FlashInfer or SGLang ground truth. Use when validating kernel definitions, adding tests for new op_types, or verifying reference implementations are correct.
# Add Reference Tests Add tests to validate reference implementations in the HuggingFace dataset clone at `tmp/flashinfer-trace/`. Ground truth is sourced from FlashInfer repository or SGLang when FlashInfer doesn't have the implementation. ## Description This skill creates test cases under `tmp/flashinfer-trace/tests/references/` (the HF dataset clone — the in-repo `flashinfer_trace/` directory was removed in the trace-dataset refactor) to validate that reference implementations in Definition JSON files produce correct outputs. The ground truth comes from: 1. **FlashInfer repository** (preferred): Official optimized GPU kernels in `tmp/flashinfer/` 2. **SGLang repository** (fallback): When FlashInfer doesn't have the kernel, use `tmp/sglang/` ## Usage ```bash # Test all definitions of a specific op_type /add-reference-tests --op-type mla_paged /add-reference-tests --op-type moe /add-reference-tests --op-type gqa_paged /add-reference-tests --op-type rmsnorm # Test a specific definition /add-reference-tests --definition-name mla_paged_decode_h16_ckv512_kpe64_ps1 # Test all definitions in the definitions directory /add-reference-tests --all # Test with custom tolerance /add-reference-tests --definition-name rmsnorm_h4096 --tolerance 1e-4 ``` ## Parameters - `definition_name` (optional): Specific definition to test (e.g., "mla_paged_decode_h16_ckv512_kpe64_ps1") - `op_type` (optional): Test all definitions of a specific op_type (e.g., "mla_paged", "moe", "rmsnorm") - `all` (optional): Test all definitions in the definitions directory - `test_sizes` (optional): List of test sizes ["small", "medium", "large"] (default: ["small", "medium"]) - `tolerance` (optional): Numerical tolerance for comparison (default: 1e-3 for fp16, 1e-5 for fp32) ## Prerequisites Run `/clone-repos` first to set up the `tmp/` directory with SGLang, FlashInfer, and the HuggingFace trace dataset clone at `tmp/flashinfer-trace/` — that clone is the only home for definitions and reference tests. ## What This Skill Does ### Phase 1: Definition Discovery 1. **Load Target Definitions**: - If `definition_name` specified: load single definition - If `op_type` specified: load all definitions matching op_type from `tmp/flashinfer-trace/definitions/{op_type}/` - If `all`: scan all definitions 2. **Check Existing Tests**: - Scan `tmp/flashinfer-trace/tests/references/` for existing test files - Skip definitions that already have tests (unless force=true) 3. **Parse Definition Schema**: - Extract axes (const/var), inputs, outputs - Identify required shapes and dtypes - Parse reference implementation code ### Phase 2: Ground Truth Discovery For each definition, locate ground truth implementation using this priority order: #### For Model Constants: HuggingFace + SGLang (Required) **Note**: See [extract-kernel-definitions](../extract-kernel-definitions/SKILL.md#for-model-constants-huggingface--sglang-required) for detailed guidance on sourcing model constants from HuggingFace and SGLang. #### For Ground Truth Execution: FlashInfer API (Primary) - **When**: FlashInfer has the kernel implementation (MOST kernels) - **Location**: `tmp/flashinfer/python/flashinfer/` - **Use for**: Running optimized GPU kernel as ground truth - **Examples**: ```python import flashinfer # GQA decode flashinfer.BatchDecodeWithPagedKVCacheWrapper(...) # GQA prefill flashinfer.BatchPrefillWithPagedKVCacheWrapper(...) # MLA flashinfer.mla.BatchMLAPagedAttentionWrapper(...) # RMSNorm flashinfer.norm.rmsnorm(...) ``` - **Important**: FlashInfer is the PRIMARY ground truth for correctness validation #### For Ground Truth Execution: SGLang (Fallback ONLY) - **When**: FlashInfer does NOT have the kernel (e.g., some MoE variants) - **Location**: `tmp/sglang/python/sglang/srt/layers/` - **Use for**: Ground truth when FlashInfer unavailable - **Examples**: ``` layers/moe/fused_moe.py # MoE when FlashInfer MoE unavailable ``` - **Important**: ONLY use SGLang as ground truth when FlashInfer doesn't support the kernel #### Ground Truth Source Mapping | Op Type | Ground Truth Source | FlashInfer API | Fallback (if FlashInfer unavailable) | |---------|---------------------|----------------|--------------------------------------| | `rmsnorm` | FlashInfer | `flashinfer.norm.rmsnorm` | N/A (FlashInfer has it) | | `fused_add_rmsnorm` | FlashInfer | `flashinfer.norm.fused_add_rmsnorm` | N/A (FlashInfer has it) | | `gqa_paged` | FlashInfer | `flashinfer.BatchDecodeWithPagedKVCacheWrapper`, `flashinfer.BatchPrefillWithPagedKVCacheWrapper` | N/A | | `gqa_ragged` | FlashInfer | `flashinfer.BatchPrefillWithRaggedKVCacheWrapper` | N/A | | `mla_paged` | FlashInfer | `flashinfer.mla.BatchMLAPagedAttentionWrapper` | N/A (FlashInfer has it) | | `moe` | SGLang (fallback) | N/A (FlashInfer MoE may not cover all variants) | `sglang/layers/moe/fused_moe.py` | | `gemm` | PyTorch | N/A | `torch.nn.functional.linear` | | `sampling` | FlashInfer | `flashinfer.sampling.*` | N/A | | `rope` | FlashInfer | `flashinfer.apply_rope_with_cos_sin_cache_inplace` | N/A | #### Reference `run()` Function Sources **Note**: For detailed guidance on sourcing reference implementations, see the [extract-kernel-definitions](../extract-kernel-definitions/SKILL.md) skill's "Reference Implementation Sources" section. **Quick Reference**: - **Primary**: FlashInfer unit tests at `tmp/flashinfer/tests/` (e.g., `test_batch_decode.py`, `test_norm.py`) - **Fallback**: SGLang vanilla implementations at `tmp/sglang/python/sglang/srt/layers/` (only when FlashInfer unavailable) ### Phase 3: Test Generation For each definition, generate test file following the standards below. ## Test File Standards ### File Structure Each test file should follow this structure: ```python import json import math from pathlib import Path import numpy as np import pytest import torch # Ground truth imports (with availability checks) try: import flashinfer from flashinfer.xxx import some_kernel FLASHINFER_AVAILABLE = True except ImportError: FLASHINFER_AVAILABLE = False # Module-level constants from definition HIDDEN_SIZE = 7168 NUM_EXPERTS = 256 # ... other constants TRACE_ROOT = Path(__file__).resolve().parents[2] WORKLOAD_JSONL_PATH = TRACE_ROOT / "workloads" / "op_type" / "definition_name.jsonl" @torch.no_grad() def run(...): """Reference implementation matching the definition.""" # Check constants assert hidden_size == HIDDEN_SIZE ... def generate_random_inputs(..., device="cuda"): """Generate random inputs for testing.""" ... return {...} def test_correctness(..., atol=1e-2, rtol=5e-2): """Test correctness of reference implementation against ground truth.""" ... def main(): """Run comprehensive tests.""" ... if __name__ == "__main__": main() ``` ### Coding Style Patterns 1. **Constants at Module Level**: Define model-specific constants at the top ```python # DeepSeek V3/R1 MoE constants HIDDEN_SIZE = 7168 INTERMEDIATE_SIZE = 2048 NUM_EXPERTS_GLOBAL = 256 NUM_LOCAL_EXPERTS = 32 # EP=8 ``` 2. **Use `@torch.no_grad()` Decorator**: For all reference implementations and test functions 3. **Input Generator Function**: Separate function `generate_random_inputs(...)` that returns a dict 4. **Test Function Pattern**: ```python def test_correctness(batch_size=4, max_seq_len=64, atol=1e-2, rtol=5e-2): """Test correctness of reference implementation against ground truth.""" print(f"\n{'='*60}") print(f"Testing {description}: {params}") print(f"{'='*60}") device = "cuda" if torch.cuda.is_available() else "cpu" if device == "cpu": print("WARNING: CUDA not available, skipping test") return # Generate inputs inputs = generate_random_inputs(...) # Run reference implementation print("\nRunning reference implementation...") ref_output = run(**inputs) # Run ground truth print("Running ground truth (FlashInfer/SGLang)...") gt_output = ground_truth_fn(**inputs) # Compare outputs print("\nComparing outputs...") # ... detailed comparison ``` ### Correctness Checking Patterns #### Standard Tolerance Check (FP16/BF16) ```python # Convert to float32 for comparison ref_f32 = ref_output.float() gt_f32 = gt_output.float() # Compute detailed error metrics abs_diff = torch.abs(ref_f32 - gt_f32) rel_diff = abs_diff / (torch.abs(gt_f32) + 1e-8) max_abs_diff = abs_diff.max().item() max_rel_diff = rel_diff.max().item() mean_abs_diff = abs_diff.mean().item() mean_rel_diff = rel_diff.mean().item() print(f"\nOutput tensor comparison:") print(f"Max absolute difference: {max_abs_diff:.6e}") print(f"Max relative difference: {max_rel_diff:.6e}") print(f"Mean absolute difference: {mean_abs_diff:.6e}") print(f"Mean relative difference: {mean_rel_diff:.6e}") # Cosine similarity and MSE cos_sim = torch.nn.functional.cosine_similarity( ref_f32.flatten(), gt_f32.flatten(), dim=0 ).item() mse = torch.mean((ref_f32 - gt_f32) ** 2).item() print(f"Cosine similarity: {cos_sim:.6f}") print(f"MSE: {mse:.6e}") # Check tolerance all_close = torch.allclose(ref_f32, gt_f32, atol=atol, rtol=rtol) if all_close: print(f"\n✓ PASSED: Outputs match within tolerance (atol={atol}, rtol={rtol})") else: print(f"\n✗ FAILED: Outputs differ beyond tolerance (atol={atol}, rtol={rtol})") ``` #### Hit Ratio Check (for FP8/Quantized Kernels) For quantized kernels with higher variance, use hit ratio instead of strict allclose: ```python # Check what percentage of elements pass the tolerance check left = (ref_f32 - gt_f32).abs() right = atol + rtol * gt_f32.abs() ok = left <= right hit_ratio = ok.float().mean().item() print(f"\nHit ratio: {hit_ratio * 100:.2f}% (need >= {percent * 100:.2f}%)") return hit_ratio >= percent # e.g., 85% ``` #### Error Location Debugging When tests fail, show top error locations: ```python if not all_close: flat = abs_diff.flatten() k = min(5, flat.numel()) topv, topi = torch.topk(flat, k) print(f"\nTop-{k} absolute error locations:") for rank in range(k): idx = topi[rank].item() # Convert flat index to multi-dimensional # ... compute indices print(f" [{indices}]: ref={ref_val:.6e}, gt={gt_val:.6e}, diff={topv[rank].item():.6e}") ``` ### Tolerance Guidelines | Data Type | atol | rtol | Notes | |-----------|------|------|-------| | float32 | 1e-5 | 1e-5 | Strictest | | float16 | 1e-3 | 1e-3 | Standard | | bfloat16 | 8e-3 | 1e-2 | 0.8% abs, 1% rel | | float8_e4m3fn | 1e-1 | 2e-1 | Use hit ratio ≥85% | | nvfp4 | 1e-1 | 2e-1 | Use hit ratio ≥85% | ## Multi-Ground-Truth Testing ### Pattern for Multiple Ground Truths When both FlashInfer and SGLang implementations are available, test against both: ```python # Ground truth imports try: from flashinfer.xxx import flashinfer_kernel FLASHINFER_AVAILABLE = True except ImportError:
عرض على GitHub
ملف SKILL.md هذا كبير جدا، لذلك يعرض SkillsMP القسم الاول فقط هنا. عرض على GitHub