mirror of https://github.com/razor-ai/soup.git
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
This commit is contained in:
parent
e30a637f48
commit
42b56f1570
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]'[/]"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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]'"
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Reference in New Issue