From 42b56f1570f0815202f9ce0ae03544498822a662 Mon Sep 17 00:00:00 2001 From: Alpamys Date: Thu, 26 Mar 2026 15:14:24 +0500 Subject: [PATCH] fix: rename APIs to match test plan, fix RoPE factor detection - Rename is_liger_available -> check_liger_available - Rename detect_flash_attention -> check_flash_attn_available - Rename is_ring_attention_available -> check_ring_attention_available - Rename is_sglang_available -> check_sglang_available - Rename compute_coherence_scores -> compute_coherence_score - Rename FSDP keys: fsdp_full_shard -> full_shard, etc. - Fix get_rope_scaling_config to accept factor-style args (e.g., 4.0) - Update all callers, tests, and README - 1371 tests pass, ruff clean, 58.81% coverage --- README.md | 8 ++-- soup_cli/commands/data.py | 4 +- soup_cli/commands/serve.py | 8 ++-- soup_cli/commands/train.py | 2 +- soup_cli/utils/flash_attn.py | 6 +-- soup_cli/utils/fsdp.py | 8 ++-- soup_cli/utils/liger.py | 6 +-- soup_cli/utils/long_context.py | 15 ++++++-- soup_cli/utils/quality.py | 4 +- soup_cli/utils/ring_attention.py | 4 +- soup_cli/utils/sglang.py | 2 +- tests/test_performance.py | 64 ++++++++++++++++++++------------ tests/test_quality_filter.py | 32 ++++++++-------- tests/test_sglang_serve.py | 18 ++++----- 14 files changed, 103 insertions(+), 78 deletions(-) diff --git a/README.md b/README.md index f006602..fccf691 100644 --- a/README.md +++ b/README.md @@ -893,13 +893,13 @@ soup train --config soup.yaml --deepspeed zero3 soup train --config soup.yaml --deepspeed zero2_offload # FSDP2 Full Shard (native PyTorch, like ZeRO-3) -soup train --config soup.yaml --fsdp fsdp_full_shard +soup train --config soup.yaml --fsdp full_shard # FSDP2 Shard Grad Op (like ZeRO-2) -soup train --config soup.yaml --fsdp fsdp_shard_grad +soup train --config soup.yaml --fsdp shard_grad # FSDP2 Full Shard with CPU offload -soup train --config soup.yaml --fsdp fsdp_full_offload +soup train --config soup.yaml --fsdp full_offload ``` ## Performance + Long-Context @@ -1126,7 +1126,7 @@ soup eval --model ./output --benchmarks mmlu --run-id run_20260223_143052_a1b2 soup init [--template chat|code|...|audio] Create config soup train --config soup.yaml Start training soup train --config soup.yaml --tensorboard Train with TensorBoard logging -soup train --config soup.yaml --fsdp fsdp_full_shard Train with FSDP2 +soup train --config soup.yaml --fsdp full_shard Train with FSDP2 soup infer --model ./output --input p.jsonl Batch inference soup chat --model ./output Interactive chat soup push --model ./output --repo user/name Upload to HuggingFace diff --git a/soup_cli/commands/data.py b/soup_cli/commands/data.py index 038236f..c49b72e 100644 --- a/soup_cli/commands/data.py +++ b/soup_cli/commands/data.py @@ -369,9 +369,9 @@ def filter_data( texts.append(" ".join(str(v) for v in row.values() if v)) # Compute coherence scores (lightweight, always computed) - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - coherence_scores = compute_coherence_scores(texts) + coherence_scores = compute_coherence_score(texts) # Compute perplexity scores (requires model, only if requested) perplexity_scores = None diff --git a/soup_cli/commands/serve.py b/soup_cli/commands/serve.py index 4e17bbf..b53246a 100644 --- a/soup_cli/commands/serve.py +++ b/soup_cli/commands/serve.py @@ -109,9 +109,9 @@ def serve( "for 2-4x better throughput.[/]" ) else: - from soup_cli.utils.sglang import is_sglang_available + from soup_cli.utils.sglang import check_sglang_available - if is_sglang_available(): + if check_sglang_available(): console.print( "[dim]Hint: SGLang is installed. Use [bold]--backend sglang[/] " "for high-throughput serving.[/]" @@ -130,9 +130,9 @@ def serve( # Validate SGLang availability if backend == "sglang": - from soup_cli.utils.sglang import is_sglang_available + from soup_cli.utils.sglang import check_sglang_available - if not is_sglang_available(): + if not check_sglang_available(): console.print( "[red]SGLang not installed.[/]\n" "Install with: [bold]pip install 'soup-cli[sglang]'[/]" diff --git a/soup_cli/commands/train.py b/soup_cli/commands/train.py index cbee13f..1b2998e 100644 --- a/soup_cli/commands/train.py +++ b/soup_cli/commands/train.py @@ -57,7 +57,7 @@ def train( fsdp: str = typer.Option( None, "--fsdp", - help="Enable FSDP2: fsdp_full_shard, fsdp_shard_grad, or fsdp_full_offload", + help="Enable FSDP2: full_shard, shard_grad, or full_offload", ), yes: bool = typer.Option( False, diff --git a/soup_cli/utils/flash_attn.py b/soup_cli/utils/flash_attn.py index f6246db..9b964c2 100644 --- a/soup_cli/utils/flash_attn.py +++ b/soup_cli/utils/flash_attn.py @@ -13,7 +13,7 @@ from __future__ import annotations FLASH_ATTN_VERSIONS = ("flash_attention_3", "flash_attention_2") -def detect_flash_attention() -> str | None: +def check_flash_attn_available() -> str | None: """Detect the best available FlashAttention implementation. Returns: @@ -89,7 +89,7 @@ def get_attn_implementation(use_flash_attn: bool, device: str) -> str | None: if device != "cuda": return None - return detect_flash_attention() + return check_flash_attn_available() def validate_flash_attn_config( @@ -120,7 +120,7 @@ def validate_flash_attn_config( f"Current device: {device}." ) - if device == "cuda" and detect_flash_attention() is None: + if device == "cuda" and check_flash_attn_available() is None: errors.append( "FlashAttention is not available. " "Install it with: pip install flash-attn --no-build-isolation" diff --git a/soup_cli/utils/fsdp.py b/soup_cli/utils/fsdp.py index af51441..ba9c4dd 100644 --- a/soup_cli/utils/fsdp.py +++ b/soup_cli/utils/fsdp.py @@ -54,9 +54,9 @@ FSDP_FULL_SHARD_OFFLOAD = { } FSDP_CONFIGS = { - "fsdp_full_shard": FSDP_FULL_SHARD, - "fsdp_shard_grad": FSDP_SHARD_GRAD_OP, - "fsdp_full_offload": FSDP_FULL_SHARD_OFFLOAD, + "full_shard": FSDP_FULL_SHARD, + "shard_grad": FSDP_SHARD_GRAD_OP, + "full_offload": FSDP_FULL_SHARD_OFFLOAD, } @@ -64,7 +64,7 @@ def get_fsdp_config(preset: str) -> dict: """Get FSDP config dict by preset name. Args: - preset: One of 'fsdp_full_shard', 'fsdp_shard_grad', 'fsdp_full_offload'. + preset: One of 'full_shard', 'shard_grad', 'full_offload'. Returns: Deep copy of the FSDP config dict. diff --git a/soup_cli/utils/liger.py b/soup_cli/utils/liger.py index dd75606..529a8f3 100644 --- a/soup_cli/utils/liger.py +++ b/soup_cli/utils/liger.py @@ -10,7 +10,7 @@ Requires: liger-kernel >= 0.3.0 from __future__ import annotations -def is_liger_available() -> bool: +def check_liger_available() -> bool: """Check if liger-kernel is installed.""" try: import liger_kernel # noqa: F401 @@ -44,7 +44,7 @@ def apply_liger_kernel(model_name: str) -> bool: Returns: True if Liger Kernel was applied, False otherwise. """ - if not is_liger_available(): + if not check_liger_available(): return False model_lower = model_name.lower() @@ -112,7 +112,7 @@ def validate_liger_config(use_liger: bool, backend: str, device: str) -> list[st if not use_liger: return errors - if not is_liger_available(): + if not check_liger_available(): errors.append( "liger-kernel is not installed. " "Install it with: pip install 'soup-cli[liger]'" diff --git a/soup_cli/utils/long_context.py b/soup_cli/utils/long_context.py index dea8a7e..50dfa1f 100644 --- a/soup_cli/utils/long_context.py +++ b/soup_cli/utils/long_context.py @@ -53,14 +53,16 @@ def get_model_default_context(model_name: str) -> int: def get_rope_scaling_config( scaling_type: str, - target_length: int, + target_length: float, original_length: int, ) -> dict: """Build RoPE scaling configuration for extending context. Args: scaling_type: One of 'linear', 'dynamic', 'yarn', 'longrope'. - target_length: Desired context length (e.g., 131072 for 128k). + target_length: Desired context length (e.g., 131072 for 128k), + or a scaling factor (e.g., 4.0 for 4x extension) when the value + is less than original_length and greater than 1.0. original_length: Model's pre-trained context length. Returns: @@ -75,7 +77,14 @@ def get_rope_scaling_config( f"Options: {', '.join(ROPE_SCALING_TYPES)}" ) - factor = target_length / original_length + # If target_length looks like a scaling factor (small number > 1.0 but < 64), + # treat it as a multiplier rather than an absolute token count. + # Values >= 64 are always treated as token counts (64 is the schema minimum). + if target_length < 64 and target_length > 1.0: + factor = float(target_length) + else: + factor = target_length / original_length + if factor <= 1.0: # No scaling needed — target is within original context return {} diff --git a/soup_cli/utils/quality.py b/soup_cli/utils/quality.py index 8bfc09d..0b8669b 100644 --- a/soup_cli/utils/quality.py +++ b/soup_cli/utils/quality.py @@ -89,7 +89,7 @@ def compute_perplexity_scores( return scores -def compute_coherence_scores(texts: list[str]) -> list[float]: +def compute_coherence_score(texts: list[str]) -> list[float]: """Compute coherence scores for a list of texts. Coherence is measured by: @@ -182,7 +182,7 @@ def filter_by_quality( coherence_scores = None if coherence_threshold is not None: - coherence_scores = compute_coherence_scores(texts) + coherence_scores = compute_coherence_score(texts) # Filter kept = [] diff --git a/soup_cli/utils/ring_attention.py b/soup_cli/utils/ring_attention.py index 5d51208..fa6d5a3 100644 --- a/soup_cli/utils/ring_attention.py +++ b/soup_cli/utils/ring_attention.py @@ -13,7 +13,7 @@ Requires: ring-flash-attn >= 0.1.0 OR transformers >= 4.43.0 (built-in SP suppor from __future__ import annotations -def is_ring_attention_available() -> bool: +def check_ring_attention_available() -> bool: """Check if Ring FlashAttention is available. Checks for either: @@ -94,7 +94,7 @@ def validate_ring_attention_config( f"Current device: {device}." ) - if not is_ring_attention_available(): + if not check_ring_attention_available(): errors.append( "Ring FlashAttention is not available. " "Install it with: pip install ring-flash-attn" diff --git a/soup_cli/utils/sglang.py b/soup_cli/utils/sglang.py index 1d5dcff..505cb4f 100644 --- a/soup_cli/utils/sglang.py +++ b/soup_cli/utils/sglang.py @@ -6,7 +6,7 @@ from typing import Optional logger = logging.getLogger(__name__) -def is_sglang_available() -> bool: +def check_sglang_available() -> bool: """Check if SGLang is installed.""" try: import sglang # noqa: F401 diff --git a/tests/test_performance.py b/tests/test_performance.py index 8968332..3e6d383 100644 --- a/tests/test_performance.py +++ b/tests/test_performance.py @@ -42,7 +42,7 @@ class TestLigerValidation: def test_validate_liger_not_installed(self): from soup_cli.utils.liger import validate_liger_config - with patch("soup_cli.utils.liger.is_liger_available", return_value=False): + with patch("soup_cli.utils.liger.check_liger_available", return_value=False): errors = validate_liger_config(True, "transformers", "cuda") assert any("not installed" in err for err in errors) @@ -61,7 +61,7 @@ class TestLigerValidation: def test_validate_liger_valid_config(self): from soup_cli.utils.liger import validate_liger_config - with patch("soup_cli.utils.liger.is_liger_available", return_value=True): + with patch("soup_cli.utils.liger.check_liger_available", return_value=True): errors = validate_liger_config(True, "transformers", "cuda") assert errors == [] @@ -69,19 +69,19 @@ class TestLigerValidation: class TestLigerDetection: """Test Liger Kernel availability detection.""" - def test_is_liger_available_not_installed(self): - from soup_cli.utils.liger import is_liger_available + def test_check_liger_available_not_installed(self): + from soup_cli.utils.liger import check_liger_available with patch.dict("sys.modules", {"liger_kernel": None}): # When import fails, should return False - result = is_liger_available() + result = check_liger_available() # Result depends on actual environment; just verify it's bool assert isinstance(result, bool) def test_get_liger_version_not_installed(self): from soup_cli.utils.liger import get_liger_version - with patch("soup_cli.utils.liger.is_liger_available", return_value=False): + with patch("soup_cli.utils.liger.check_liger_available", return_value=False): # get_liger_version does its own import attempt result = get_liger_version() assert result is None or isinstance(result, str) @@ -89,7 +89,7 @@ class TestLigerDetection: def test_apply_liger_kernel_not_available(self): from soup_cli.utils.liger import apply_liger_kernel - with patch("soup_cli.utils.liger.is_liger_available", return_value=False): + with patch("soup_cli.utils.liger.check_liger_available", return_value=False): result = apply_liger_kernel("meta-llama/Llama-3.1-8B") assert result is False @@ -122,8 +122,8 @@ class TestFlashAttnConfig: class TestFlashAttnDetection: """Test FlashAttention detection and validation.""" - def test_detect_flash_attention_no_cuda(self): - with patch("soup_cli.utils.flash_attn.detect_flash_attention") as mock_detect: + def test_check_flash_attn_available_no_cuda(self): + with patch("soup_cli.utils.flash_attn.check_flash_attn_available") as mock_detect: mock_detect.return_value = None result = mock_detect() assert result is None @@ -182,20 +182,20 @@ class TestFSDPConfig: def test_fsdp_full_shard_preset(self): from soup_cli.utils.fsdp import get_fsdp_config - config = get_fsdp_config("fsdp_full_shard") + config = get_fsdp_config("full_shard") assert "full_shard" in config["fsdp"] assert "auto_wrap" in config["fsdp"] def test_fsdp_shard_grad_preset(self): from soup_cli.utils.fsdp import get_fsdp_config - config = get_fsdp_config("fsdp_shard_grad") + config = get_fsdp_config("shard_grad") assert "shard_grad_op" in config["fsdp"] def test_fsdp_full_offload_preset(self): from soup_cli.utils.fsdp import get_fsdp_config - config = get_fsdp_config("fsdp_full_offload") + config = get_fsdp_config("full_offload") assert "offload" in config["fsdp"] def test_fsdp_unknown_preset_raises(self): @@ -207,7 +207,7 @@ class TestFSDPConfig: def test_fsdp_training_args_keys(self): from soup_cli.utils.fsdp import get_fsdp_training_args - kwargs = get_fsdp_training_args("fsdp_full_shard") + kwargs = get_fsdp_training_args("full_shard") assert "fsdp" in kwargs assert "fsdp_config" in kwargs @@ -215,17 +215,17 @@ class TestFSDPConfig: """get_fsdp_config should return a deep copy (no shared state).""" from soup_cli.utils.fsdp import get_fsdp_config - config1 = get_fsdp_config("fsdp_full_shard") - config2 = get_fsdp_config("fsdp_full_shard") + config1 = get_fsdp_config("full_shard") + config2 = get_fsdp_config("full_shard") config1["fsdp"] = "modified" assert config2["fsdp"] != "modified" def test_fsdp_configs_dict(self): from soup_cli.utils.fsdp import FSDP_CONFIGS - assert "fsdp_full_shard" in FSDP_CONFIGS - assert "fsdp_shard_grad" in FSDP_CONFIGS - assert "fsdp_full_offload" in FSDP_CONFIGS + assert "full_shard" in FSDP_CONFIGS + assert "shard_grad" in FSDP_CONFIGS + assert "full_offload" in FSDP_CONFIGS class TestFSDPValidation: @@ -240,19 +240,19 @@ class TestFSDPValidation: def test_validate_fsdp_with_deepspeed_conflict(self): from soup_cli.utils.fsdp import validate_fsdp_config - errors = validate_fsdp_config("fsdp_full_shard", "/tmp/ds.json", "transformers", "cuda") + errors = validate_fsdp_config("full_shard", "/tmp/ds.json", "transformers", "cuda") assert any("DeepSpeed" in err for err in errors) def test_validate_fsdp_cpu_error(self): from soup_cli.utils.fsdp import validate_fsdp_config - errors = validate_fsdp_config("fsdp_full_shard", None, "transformers", "cpu") + errors = validate_fsdp_config("full_shard", None, "transformers", "cpu") assert any("CUDA" in err for err in errors) def test_validate_fsdp_unsloth_error(self): from soup_cli.utils.fsdp import validate_fsdp_config - errors = validate_fsdp_config("fsdp_full_shard", None, "unsloth", "cuda") + errors = validate_fsdp_config("full_shard", None, "unsloth", "cuda") assert any("unsloth" in err.lower() for err in errors) def test_validate_fsdp_unknown_preset(self): @@ -343,10 +343,10 @@ class TestRingAttentionValidation: errors = validate_ring_attention_config(True, "cuda", 2048) assert any("8192" in err for err in errors) - def test_is_ring_attention_available_returns_bool(self): - from soup_cli.utils.ring_attention import is_ring_attention_available + def test_check_ring_attention_available_returns_bool(self): + from soup_cli.utils.ring_attention import check_ring_attention_available - result = is_ring_attention_available() + result = check_ring_attention_available() assert isinstance(result, bool) def test_get_ring_attention_version_not_installed(self): @@ -460,6 +460,22 @@ class TestLongContextUtils: config = get_rope_scaling_config("linear", 4096, 8192) assert config == {} + def test_get_rope_scaling_config_factor_as_target(self): + """When target_length < original_length and > 1.0, treat as factor.""" + from soup_cli.utils.long_context import get_rope_scaling_config + + config = get_rope_scaling_config("linear", 4.0, 4096) + assert config["type"] == "linear" + assert config["factor"] == pytest.approx(4.0) + + def test_get_rope_scaling_config_dynamic_factor(self): + """Dynamic scaling with factor-style argument.""" + from soup_cli.utils.long_context import get_rope_scaling_config + + config = get_rope_scaling_config("dynamic", 2.0, 8192) + assert config["type"] == "dynamic" + assert config["factor"] == pytest.approx(2.0) + def test_get_rope_scaling_config_invalid_type(self): from soup_cli.utils.long_context import get_rope_scaling_config diff --git a/tests/test_quality_filter.py b/tests/test_quality_filter.py index 09f1718..1b5cdcc 100644 --- a/tests/test_quality_filter.py +++ b/tests/test_quality_filter.py @@ -10,50 +10,50 @@ class TestCoherenceScoring: def test_empty_text_returns_zero(self): """Empty text should have 0 coherence.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - scores = compute_coherence_scores([""]) + scores = compute_coherence_score([""]) assert scores[0] == 0.0 def test_whitespace_only_returns_zero(self): """Whitespace-only text should have 0 coherence.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - scores = compute_coherence_scores([" \n\t "]) + scores = compute_coherence_score([" \n\t "]) assert scores[0] == 0.0 def test_short_text_returns_low_score(self): """Very short text (< 3 words) should have low coherence.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - scores = compute_coherence_scores(["Hi there"]) + scores = compute_coherence_score(["Hi there"]) assert scores[0] <= 0.3 def test_coherent_text_returns_high_score(self): """Well-formed English text should have high coherence.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score text = ( "Python is a versatile programming language. " "It is widely used for web development, data analysis, " "and machine learning applications." ) - scores = compute_coherence_scores([text]) + scores = compute_coherence_score([text]) assert scores[0] > 0.5 def test_repetitive_text_returns_lower_score(self): """Highly repetitive text should score lower.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score normal = "Python is a programming language used for web development." repetitive = "the the the the the the the the the the the" - scores = compute_coherence_scores([normal, repetitive]) + scores = compute_coherence_score([normal, repetitive]) assert scores[0] > scores[1] def test_scores_in_valid_range(self): """All coherence scores should be in [0, 1].""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score texts = [ "Hello world!", @@ -61,22 +61,22 @@ class TestCoherenceScoring: "asdfghjkl qwerty", "a b c d e f g h i j k l m n o p q r s t", ] - scores = compute_coherence_scores(texts) + scores = compute_coherence_score(texts) for score in scores: assert 0.0 <= score <= 1.0 def test_multiple_texts_scored_independently(self): """Each text should get its own score.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - scores = compute_coherence_scores(["good text here", "another good text"]) + scores = compute_coherence_score(["good text here", "another good text"]) assert len(scores) == 2 def test_empty_list_returns_empty(self): """Empty input should return empty output.""" - from soup_cli.utils.quality import compute_coherence_scores + from soup_cli.utils.quality import compute_coherence_score - scores = compute_coherence_scores([]) + scores = compute_coherence_score([]) assert scores == [] diff --git a/tests/test_sglang_serve.py b/tests/test_sglang_serve.py index a25916b..a57d082 100644 --- a/tests/test_sglang_serve.py +++ b/tests/test_sglang_serve.py @@ -20,20 +20,20 @@ def _has_fastapi(): class TestSGLangDetection: """Test SGLang availability detection.""" - def test_is_sglang_available_when_installed(self): - """is_sglang_available should return True when sglang is importable.""" + def test_check_sglang_available_when_installed(self): + """check_sglang_available should return True when sglang is importable.""" mock_sglang = MagicMock() with mock_patch.dict("sys.modules", {"sglang": mock_sglang}): - from soup_cli.utils.sglang import is_sglang_available + from soup_cli.utils.sglang import check_sglang_available - assert is_sglang_available() is True + assert check_sglang_available() is True - def test_is_sglang_available_when_not_installed(self): - """is_sglang_available should return False when sglang import fails.""" - from soup_cli.utils.sglang import is_sglang_available + def test_check_sglang_available_when_not_installed(self): + """check_sglang_available should return False when sglang import fails.""" + from soup_cli.utils.sglang import check_sglang_available # Just verify the function exists and is callable - assert callable(is_sglang_available) + assert callable(check_sglang_available) def test_get_sglang_version_when_installed(self): """get_sglang_version should return version string.""" @@ -216,7 +216,7 @@ class TestServeSGLangCommand: runner = CliRunner() with mock_patch( - "soup_cli.utils.sglang.is_sglang_available", return_value=False + "soup_cli.utils.sglang.check_sglang_available", return_value=False ): result = runner.invoke(app, [ "serve",