Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
85 changes: 85 additions & 0 deletions config/experiments.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -341,3 +341,88 @@ experiments:
- rouge2
- rougeL


# QMSum - DeepSeek V4 Flash baseline
- id: qmsum_baseline_deepseekv4flash
tags: [qmsum]
model:
name: deepseek:deepseek-v4-flash
mode: baseline
generation:
max_tokens: 256
temperature: 1.2
top_p: 1.0
dataset:
name: qmsum
split: test
num_examples: 281
seed: 40
task:
name: summarization
metrics:
- rouge1
- rouge2
- rougeL

# QMSum - DeepSeek V4 Pro baseline
- id: qmsum_baseline_deepseekv4pro
tags: [qmsum]
model:
name: deepseek:deepseek-v4-pro
mode: baseline
generation:
max_tokens: 256
temperature: 1.2
top_p: 1.0
dataset:
name: qmsum
split: test
num_examples: 281
seed: 40
task:
name: summarization
metrics:
- rouge1
- rouge2
- rougeL

# CLINC OOS - DeepSeek V4 Flash baseline (Classify task)
# NOTE: num_examples/seed here are placeholders — confirm the sample size
# and seed used to produce the other models' Classify numbers (not found
# anywhere in this repo) before treating this as comparable.
- id: classify_clinc_baseline_deepseekv4flash
tags: [clinc_oos]
model:
name: deepseek:deepseek-v4-flash
mode: baseline
generation:
max_tokens: 16
temperature: 0.0
dataset:
name: clinc_oos
split: test
num_examples: 400
seed: 42
task:
name: classification
metrics:
- exact_match

# CLINC OOS - DeepSeek V4 Pro baseline (Classify task)
- id: classify_clinc_baseline_deepseekv4pro
tags: [clinc_oos]
model:
name: deepseek:deepseek-v4-pro
mode: baseline
generation:
max_tokens: 16
temperature: 0.0
dataset:
name: clinc_oos
split: test
num_examples: 400
seed: 42
task:
name: classification
metrics:
- exact_match
6 changes: 6 additions & 0 deletions src/core/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
# Import model clients
from models.openai_client import OpenAIClient
from models.gemini_client import GeminiClient
from models.deepseek_client import DeepSeekClient
from models.scaledown_client import ScaleDownClient
from models.scaledown_summarize_client import ScaleDownSummarizeClient
from models.base import ModelClient, PricingInfo, DEFAULT_MODEL_PRICING
Expand All @@ -16,12 +17,14 @@
from dataset.msmarco import MSMARCODataset
from dataset.financebench import FinanceBenchDataset
from dataset.qmsum import QMSumDataset
from dataset.clinc_oos import ClincOosDataset
from dataset.base import Dataset

# Import tasks
from tasks.rag_task import RAGTask
from tasks.retrieval_task import RetrievalTask
from tasks.summarization_task import SummarizationTask
from tasks.classification_task import ClassificationTask
from tasks.base import Task

# Import metrics
Expand Down Expand Up @@ -52,6 +55,7 @@
MODEL_REGISTRY: Dict[str, type] = {
"openai": OpenAIClient,
"gemini": GeminiClient,
"deepseek": DeepSeekClient,
}

DATASET_REGISTRY: Dict[str, type] = {
Expand All @@ -61,12 +65,14 @@
"msmarco": MSMARCODataset,
"financebench": FinanceBenchDataset,
"qmsum": QMSumDataset,
"clinc_oos": ClincOosDataset,
}

TASK_REGISTRY: Dict[str, type] = {
"rag_qa": RAGTask,
"retrieval_task": RetrievalTask,
"summarization": SummarizationTask,
"classification": ClassificationTask,
}

METRIC_REGISTRY: Dict[str, Callable] = {
Expand Down
61 changes: 61 additions & 0 deletions src/dataset/clinc_oos.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
"""CLINC OOS dataset loader.

CLINC OOS (Larson et al., 2019) is an intent classification benchmark with
150 in-scope intents across 10 domains plus an out-of-scope class. Used here
for the "Classify" task on the ScaleDown benchmark pages.

Loaded from HuggingFace clinc_oos ("plus" config, which includes out-of-scope
examples in addition to the 150 in-domain intents).
"""
import random
from typing import Iterator, Optional

import datasets as hf_datasets

from dataset.base import Dataset, Example


class ClincOosDataset(Dataset):
"""CLINC OOS single-turn intent classification dataset."""

def __init__(self, name: str = "clinc_oos"):
super().__init__(name)

def load_examples(
self,
split: str = "test",
limit: Optional[int] = None,
seed: int = 42,
) -> Iterator[Example]:
"""Load and yield CLINC OOS examples.

Args:
split: Dataset split (train/validation/test).
limit: Maximum number of examples to load (None = all).
seed: Random seed for deterministic sampling.

Yields:
Example dicts with {id, context ("", unused), question (query text),
answer ([intent label]), labels (full label list, for prompt building)}.
"""
hf_dataset = hf_datasets.load_dataset("clinc_oos", "plus", split=split)
label_names = hf_dataset.features["intent"].names

all_examples = list(hf_dataset)
total_size = len(all_examples)
if limit is not None and limit < total_size:
random.seed(seed)
indices = random.sample(range(total_size), limit)
indices.sort()
else:
indices = range(total_size)

for idx in indices:
row = all_examples[idx]
yield Example(
id=str(idx),
context="",
question=row["text"],
answer=[label_names[row["intent"]]],
labels=label_names,
)
4 changes: 4 additions & 0 deletions src/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,10 @@ def calculate_cost(self, input_tokens: int, output_tokens: int) -> Optional[floa
"gemini-2.5-flash": PricingInfo(input_per_1m_tokens=0.30, output_per_1m_tokens=2.50),
"gemini-2.5-flash-lite": PricingInfo(input_per_1m_tokens=0.10, output_per_1m_tokens=0.40),
"gemini-2.5-pro": PricingInfo(input_per_1m_tokens=1.25, output_per_1m_tokens=10.00),

# DeepSeek models (cache-miss input rate; see deepseek.ai/pricing, Aug 2026)
"deepseek-v4-flash": PricingInfo(input_per_1m_tokens=0.14, output_per_1m_tokens=0.28),
"deepseek-v4-pro": PricingInfo(input_per_1m_tokens=0.435, output_per_1m_tokens=0.87),
}


Expand Down
161 changes: 161 additions & 0 deletions src/models/deepseek_client.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
"""DeepSeek model client.

DeepSeek exposes an OpenAI-compatible Chat Completions API, so this client
reuses the `openai` SDK pointed at DeepSeek's base URL rather than a bespoke
HTTP layer. Token usage for cost calculation comes straight from the API
response (DeepSeek has no public tokenizer package for local counting).

Environment Variables:
DEEPSEEK_API_KEY: Required API key for DeepSeek.

Model names (verify against your account before relying on these — DeepSeek
was mid-rollout on V4 Pro as of Aug 2026, see PR description):
deepseek-v4-flash
deepseek-v4-pro
"""
import logging
import time
from typing import Optional

from openai import OpenAI, RateLimitError, APIError, APIConnectionError

from models.base import ModelClient, GenerationInput, ModelOutput, PricingInfo

logger = logging.getLogger(__name__)

DEEPSEEK_BASE_URL = "https://api.deepseek.com"


class DeepSeekClient(ModelClient):
"""Client for DeepSeek models via the OpenAI-compatible Chat Completions API."""

def __init__(
self,
model_name: str,
api_key: Optional[str] = None,
pricing: Optional[PricingInfo] = None,
max_retries: int = 3,
retry_delay: float = 2.0,
):
"""Initialize DeepSeek client.

Args:
model_name: DeepSeek model name (e.g., "deepseek-v4-flash", "deepseek-v4-pro").
api_key: DeepSeek API key. If None, uses DEEPSEEK_API_KEY env var.
pricing: Optional pricing information for cost calculation.
max_retries: Maximum number of retries for failed requests.
retry_delay: Initial delay in seconds for exponential backoff.
"""
super().__init__(model_name, api_key, pricing)
self.provider = "deepseek"
self.max_retries = max_retries
self.retry_delay = retry_delay

# DEEPSEEK_API_KEY is read by the OpenAI SDK only if api_key is passed
# explicitly here; there is no implicit env var fallback like OPENAI_API_KEY.
import os
resolved_key = api_key or os.environ.get("DEEPSEEK_API_KEY")
if not resolved_key:
raise ValueError(
"DeepSeek API key required: pass api_key or set DEEPSEEK_API_KEY."
)
self.client = OpenAI(api_key=resolved_key, base_url=DEEPSEEK_BASE_URL)

def generate(self, inp: GenerationInput) -> ModelOutput:
"""Generate text using DeepSeek's Chat Completions API.

Args:
inp: GenerationInput with prompt and configuration.

Returns:
ModelOutput with generated text and metadata.
"""
messages = []
if inp.system_prompt:
messages.append({"role": "system", "content": inp.system_prompt})

user_content = inp.user_prompt
if inp.context:
user_content = f"{inp.context}\n\n{inp.user_prompt}"
messages.append({"role": "user", "content": user_content})

params = {
"model": self.model_name,
"messages": messages,
}
if inp.temperature is not None:
params["temperature"] = inp.temperature
if inp.max_output_tokens is not None:
params["max_tokens"] = inp.max_output_tokens
if inp.top_p is not None:
params["top_p"] = inp.top_p
if inp.stop is not None:
params["stop"] = inp.stop
if inp.response_schema is not None:
# DeepSeek supports {"type": "json_object"} but not strict json_schema
# mode as of Aug 2026 — fall back to json_object and rely on the
# system prompt to describe the schema.
params["response_format"] = {"type": "json_object"}

last_exception = None
for attempt in range(self.max_retries):
try:
start_time = time.time()
response = self.client.chat.completions.create(**params)
latency_ms = (time.time() - start_time) * 1000

text = response.choices[0].message.content or ""
input_tokens = response.usage.prompt_tokens if response.usage else 0
output_tokens = response.usage.completion_tokens if response.usage else 0
cost_usd = self._calculate_cost(input_tokens, output_tokens)

return ModelOutput(
text=text,
latency_ms=latency_ms,
input_tokens=input_tokens,
output_tokens=output_tokens,
cost_usd=cost_usd,
input_question=inp.user_prompt,
input_context=inp.context,
input_system_prompt=inp.system_prompt,
)

except RateLimitError as e:
last_exception = e
if attempt < self.max_retries - 1:
delay = self.retry_delay * (2 ** attempt)
logger.info(f"Rate limit hit, retrying in {delay}s (attempt {attempt + 1}/{self.max_retries})")
time.sleep(delay)
else:
logger.error(f"Rate limit exceeded after {self.max_retries} retries")

except (APIError, APIConnectionError) as e:
last_exception = e
if attempt < self.max_retries - 1:
delay = self.retry_delay * (2 ** attempt)
logger.info(f"API error, retrying in {delay}s (attempt {attempt + 1}/{self.max_retries}): {e}")
time.sleep(delay)
else:
logger.error(f"API error after {self.max_retries} retries: {e}")
except Exception as e:
logger.error(f"Unexpected error in DeepSeekClient.generate: {e}", exc_info=True)
raise
if last_exception is not None:
raise last_exception
else:
raise RuntimeError("DeepSeekClient.generate failed with an unexpected error and no exception was captured.")

def count_tokens(self, text: str) -> int:
"""Approximate token count for text not yet sent to the API.

DeepSeek has no public local tokenizer, so this is a rough heuristic
(chars / 4) used only for pre-flight estimates. Actual cost accounting
always uses the `usage` block returned by the API in generate().

Args:
text: Input text to count tokens for.

Returns:
Approximate token count.
"""
return max(1, len(text) // 4)
Loading