timeagent / code /OpenTSLM /demo /huggingface /05_test_hf_ecg_qa_cot.py
roh8exe's picture
Upload folder using huggingface_hub
60b21d3 verified
Raw
History Blame Contribute Delete
2.82 kB
# SPDX-FileCopyrightText: 2025 Stanford University, ETH Zurich, and the project authors (see CONTRIBUTORS.md)
# SPDX-FileCopyrightText: 2025 This source file is part of the OpenTSLM open-source project.
#
# SPDX-License-Identifier: MIT
"""
Demo script for testing ECG QA CoT (ECG Question Answering Chain-of-Thought) model from HuggingFace.
This script:
1. Loads a pretrained model from HuggingFace Hub
2. Loads the ECG QA CoT test dataset
3. Generates predictions on the evaluation set
4. Prints model outputs
"""
from opentslm.model.llm.OpenTSLM import OpenTSLM
from opentslm.time_series_datasets.ecg_qa.ECGQACoTQADataset import ECGQACoTQADataset
from opentslm.time_series_datasets.util import extend_time_series_to_match_patch_size_and_aggregate
from torch.utils.data import DataLoader
from opentslm.model_config import PATCH_SIZE
import torch
# Model repository ID - change this to test different models
REPO_ID = "OpenTSLM/llama-3.2-1b-ecg-sp"
def main():
print("=" * 60)
print("ECG QA CoT Model Demo")
print("=" * 60)
# Load model from HuggingFace
print(f"\n๐Ÿ“ฅ Loading model from {REPO_ID}...")
enable_lora = False
if "-sp" in REPO_ID:
enable_lora = True
model = OpenTSLM.load_pretrained(REPO_ID, enable_lora=enable_lora, device="cuda" if torch.cuda.is_available() else "cpu")
# Create dataset
print("\n๐Ÿ“Š Loading ECG QA CoT test dataset...")
test_dataset = ECGQACoTQADataset("test", EOS_TOKEN=model.get_eos_token())
# Create data loader
test_loader = DataLoader(
test_dataset,
shuffle=False,
batch_size=1,
collate_fn=lambda batch: extend_time_series_to_match_patch_size_and_aggregate(
batch, patch_size=PATCH_SIZE
),
)
print(f"\n๐Ÿ” Running inference on {len(test_dataset)} test samples...")
print("=" * 60)
# Iterate over evaluation set
for i, batch in enumerate(test_loader):
# Generate predictions
predictions = model.generate(batch, max_new_tokens=500)
# Print results
for sample, pred in zip(batch, predictions):
print(f"\n๐Ÿ“ Sample {i + 1}:")
if 'pre_prompt' in sample:
print(f" Question: {sample['pre_prompt']}")
if 'template_id' in sample:
print(f" Template ID: {sample['template_id']}")
if 'ecg_id' in sample:
print(f" ECG ID: {sample['ecg_id']}")
print(f" Gold Answer: {sample.get('answer', 'N/A')}")
print(f" Model Output: {pred}")
print("-" * 60)
# Limit to first 5 samples for demo
if i >= 4:
print("\nโœ… Demo complete! (Showing first 5 samples)")
break
if __name__ == "__main__":
main()