| |
| |
| |
| |
|
|
| import re |
| import sys |
| from typing import Dict, Any |
|
|
| from common_evaluator import CommonEvaluator |
| from opentslm.time_series_datasets.TSQADataset import TSQADataset |
|
|
|
|
| def evaluate_tsqa(ground_truth: str, prediction: str) -> Dict[str, Any]: |
| """ |
| Evaluate TSQA predictions against ground truth. |
| |
| Args: |
| ground_truth: The correct answer |
| prediction: The model's prediction |
| |
| Returns: |
| Dictionary containing evaluation metrics |
| """ |
| |
| gt_clean = ground_truth.lower().strip() |
| pred_clean = prediction.lower().strip() |
| |
| |
| |
| |
| gt_clean, pred_clean = gt_clean[:3], pred_clean[:3] |
|
|
| |
| answer_match = re.search(r'answer:\s*(.+)', pred_clean, re.IGNORECASE) |
| if answer_match: |
| pred_answer = answer_match.group(1).strip() |
| else: |
| |
| pred_answer = pred_clean |
| |
| |
| accuracy = int(gt_clean == pred_answer) |
| |
| return { |
| "accuracy": accuracy, |
| } |
|
|
|
|
| def main(): |
| """Main function to run TSQA evaluation.""" |
| |
| if len(sys.argv) != 2: |
| print("Usage: python evaluate_tsqa.py <model_name>") |
| print("Example: python evaluate_tsqa.py meta-llama/Llama-3.2-1B") |
| sys.exit(1) |
| |
| model_name = sys.argv[1] |
| |
| |
| dataset_classes = [TSQADataset] |
| |
| |
| evaluation_functions = { |
| "TSQADataset": evaluate_tsqa, |
| } |
| |
| |
| evaluator = CommonEvaluator() |
| |
| |
| results_df = evaluator.evaluate_multiple_models( |
| model_names=[model_name], |
| dataset_classes=dataset_classes, |
| evaluation_functions=evaluation_functions, |
| max_samples=None, |
| max_new_tokens=40, |
| ) |
| |
| print("\n" + "="*80) |
| print("FINAL RESULTS SUMMARY") |
| print("="*80) |
| print(results_df.to_string(index=False)) |
| |
| return results_df |
|
|
|
|
| if __name__ == "__main__": |
| main() |