Files
OpenViking/benchmark/RAG/run.py
T
zgy 1b3a8f20a0 Feat(benchmark): Add benchmark/RAG : RAG system evaluation framework (#825)
* Add RAGbenchmark: RAG system evaluation framework

* Update README.md

* Update README.md

* Update README.md

* Code structure refactoring

* feat: improve RAG benchmark with dataset sampling and configuration updates

- Add complete dataset sampling scripts with document-level sampling
- Implement filtering logic consistent with adapters (exclude category 5 for Locomo, no answer for SyllabusQA, unanswerable for Qasper)
- Update configuration from raw_data/dataset_dir to dataset_path for clarity
- Enhance adapters with improved path handling and data loading
- Add gitignore for data and output directories
- Add dependencies (datasets, pandas, tavily-python)
- Add test files and documentation

* feat: add stratified sampling support to all datasets

- Implement stratified sampling for Locomo (by category 1-4)
- Implement stratified sampling for SyllabusQA (by question_type)
- Implement stratified sampling for Qasper (by answer type: extractive/free_form/yes_no)
- Implement stratified sampling for FinanceBench (by question_type)
- Add proper handling when sample size cannot be evenly split:
  - Display warning message
  - Distribute remaining QAs to first N categories
  - Fall back to random sampling if sample size too small
- Update prepare_dataset.py to support both 'random' and 'stratified' modes
- Set default sampling mode to 'random'

* Update locomo adapter to support image attachments and other improvements

* Update dataset documentation with actual document counts

* Add benchmark results reference and reproduction steps

* Improve sampling scripts for benchmark reproducibility

* Refactor sample_dataset.py: extract common sampling logic

- Fix two bugs:
  1. num_docs + sample_size + random path: use int indices instead of dict tuples
  2. pure stratified path: use len() for list length calculation

- Extract common sampling utilities:
  - calculate_category_targets()
  - stratified_sample_with_reallocation()
  - random_sample_qas()
  - sample_docs_stratified()
  - sample_docs_random()

- Reduce code duplication by ~60-70%
- Improve maintainability and readability
- Keep full backward compatibility

* Update config.yaml: improve configuration structure

- Add FinanceBench to supported datasets list
- Change to template configuration format
- Add execution: section for better organization

* Fix bug: duplicate worker_end() call in generation failure path

- Remove duplicate monitor.worker_end(success=False) call in run_generation()
- The _process_generation_task() already calls worker_end() in its exception handler
- This prevents double-counting of failed tasks and distorted statistics

* Fix bug: _get_required_syllabi() doesn't support JSON input

- Add JSON file support to _get_required_syllabi()
- Extract syllabus names from JSON keys (same format as _load_from_json())
- This ensures data_prepare() processes correct docx files when using JSON input

* Improve exception re-raising: use bare raise to preserve traceback

- Replace 'raise e' with bare 'raise' to preserve original traceback
- Also remove unused 'e' variable since we don't need it
- This makes debugging easier by showing where the exception actually occurred

* Fix bug: Locomo prompt uses raw gold_answer instead of gold_answer_str

- In Locomo prompt, use gold_answer_str instead of gold_answer
- This ensures consistent formatting when gold_answer is a list
- Both Locomo and Generic prompts now use the same ' | ' separated format

* Improve directory ingest: use os.path.commonpath() for robustness

- Replace manual common ancestor calculation with os.path.commonpath()
- os.path.commonpath() handles all OS path separators correctly
- Add try-except to handle ValueError when no common path exists
- More robust than manual split(os.sep) approach

* benchmark: honor skip_ingestion and fail on LLM retry exhaustion
2026-04-01 14:39:53 +08:00

170 lines
6.0 KiB
Python

import os
import sys
import yaml
import importlib
from argparse import ArgumentParser
from pathlib import Path
sys.path.append(str(Path(__file__).parent))
from src.core.logger import setup_logging
# ==========================================
# 1. Environment Initialization
# ==========================================
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = SCRIPT_DIR
ov_config_path = os.path.join(SCRIPT_DIR, "ov.conf")
if os.path.exists(ov_config_path):
os.environ["OPENVIKING_CONFIG_FILE"] = ov_config_path
print(f"[Init] Auto-detected OpenViking config: {ov_config_path}")
try:
from src.pipeline import BenchmarkPipeline
from src.core.vector_store import VikingStoreWrapper
from src.core.llm_client import LLMClientWrapper
except SyntaxError as e:
print(f"\n[Fatal Error] Syntax error while importing modules: {e}")
sys.exit(1)
except ImportError as e:
print(f"\n[Fatal Error] Cannot import modules: {e}")
print(f"Current sys.path: {sys.path}\n")
sys.exit(1)
# ==========================================
# 2. Helper Functions
# ==========================================
def load_config(config_path):
if not os.path.exists(config_path):
raise FileNotFoundError(f"Config file not found: {config_path}")
with open(config_path, 'r', encoding='utf-8') as f:
return yaml.safe_load(f)
def resolve_path(path_str, base_path):
"""
Convert relative path to absolute path based on base_path.
If path_str is already absolute, keep it unchanged.
"""
if not path_str:
return path_str
if os.path.isabs(path_str):
return path_str
return os.path.normpath(os.path.join(base_path, path_str))
# ==========================================
# 3. Main Program
# ==========================================
def main():
parser = ArgumentParser(description="Run RAG Benchmark (Smart Path Handling)")
default_config_path = os.path.join(SCRIPT_DIR, "config/config.yaml")
parser.add_argument("--config", default=default_config_path,
help=f"Path to config file. Default: {default_config_path}")
parser.add_argument("--step", choices=["all", "gen", "eval", "del"], default="all",
help="Execution step: 'gen' (Retrieval+LLM), 'eval' (Judge), or 'all'")
args = parser.parse_args()
# --- B. Load and Parse Config ---
config_path = os.path.abspath(args.config)
print(f"[Init] Loading configuration from: {config_path}")
try:
config = load_config(config_path)
except FileNotFoundError as e:
print(f"[Error] {e}")
return
# --- C. Path Resolution ---
print(f"[Init] Resolving paths relative to Project Root: {PROJECT_ROOT}")
dataset_name = config.get('dataset_name', 'UnknownDataset')
retrieval_topk = config.get('execution', {}).get('retrieval_topk', 5)
format_vars = {
'dataset_name': dataset_name,
'retrieval_topk': retrieval_topk
}
path_keys = ['dataset_path', 'output_dir', 'vector_store', 'log_file', 'doc_output_dir']
for key in path_keys:
if key in config.get('paths', {}):
original = config['paths'][key]
rendered_path = original.format(**format_vars)
resolved = resolve_path(rendered_path, PROJECT_ROOT)
config['paths'][key] = resolved
# print(f" - {key}: {resolved}")
# --- D. Initialize Components ---
try:
logger = setup_logging(config['paths']['log_file'])
logger.info(">>> Benchmark Session Started")
# 1. Adapter (Dynamic Loading)
adapter_cfg = config.get('adapter', {})
module_path = adapter_cfg.get('module', 'src.adapters.locomo_adapter')
class_name = adapter_cfg.get('class_name', 'LocomoAdapter')
logger.info(f"Dynamically loading Adapter: {class_name} from {module_path}")
logger.info(f"Loading dataset from: {config['paths']['dataset_path']}")
try:
mod = importlib.import_module(module_path)
AdapterClass = getattr(mod, class_name)
adapter = AdapterClass(raw_file_path=config['paths']['dataset_path'])
except ImportError as e:
logger.error(f"Could not import module '{module_path}'. Please check your config 'adapter.module'. Error: {e}")
raise e
except AttributeError as e:
logger.error(f"Class '{class_name}' not found in module '{module_path}'. Please check your config 'adapter.class_name'. Error: {e}")
raise e
# 2. Vector Store
vector_store = VikingStoreWrapper(store_path=config['paths']['vector_store'])
# 3. LLM Client
api_key = os.environ.get(
config['llm'].get('api_key_env_var', ''),
config['llm'].get('api_key')
)
if not api_key:
logger.warning("No API Key found in config or environment variables!")
llm_client = LLMClientWrapper(config=config['llm'], api_key=api_key)
# 4. Pipeline
pipeline = BenchmarkPipeline(
config=config,
adapter=adapter,
vector_db=vector_store,
llm=llm_client
)
# --- E. Execute Tasks ---
if args.step in ["all", "gen"]:
logger.info("Stage: Generation (Ingest -> Retrieve -> Generate)")
pipeline.run_generation()
if args.step in ["all", "eval"]:
logger.info("Stage: Evaluation (Judge -> Metrics)")
pipeline.run_evaluation()
if args.step in ["all", "del"]:
logger.info("Stage: Delete Vector Store")
pipeline.run_deletion()
logger.info("Benchmark finished successfully.")
except KeyboardInterrupt:
print("\n[Stop] Execution interrupted by user.")
except Exception as e:
if 'logger' in locals():
logger.exception("Fatal error during execution")
print(f"\n[Fatal Error] Program execution error: {str(e)}")
sys.exit(1)
if __name__ == "__main__":
main()