"""
Quick wrapper to ensure training runs on GPU with proper resource management.
"""

import torch
import subprocess
import sys
import os

# Force GPU usage
os.environ['CUDA_VISIBLE_DEVICES'] = '0'

# Verify GPU
if not torch.cuda.is_available():
    print("ERROR: GPU not available!")
    print("Available devices:", torch.device('cpu'))
    sys.exit(1)

print("="*70)
print("GPU TRAINING LAUNCHER")
print("="*70)
print(f"GPU Device: {torch.cuda.get_device_name(0)}")
print(f"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB")
print(f"CUDA Version: {torch.version.cuda}")
print("="*70)

# Clear cache
torch.cuda.empty_cache()

# Run training
print("\nStarting training on GPU...\n")
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from training.train_with_eval import main

if __name__ == '__main__':
    # Override sys.argv to set training parameters
    import sys
    if len(sys.argv) == 1:
        # Default arguments
        sys.argv = [
            'train_on_gpu.py',
            '--train_data', './data/sample_train.txt',
            '--val_data', './data/sample_valid.txt',
            '--test_data', './data/sample_test.txt',
            '--test_names', 'sample_test',
            '--epochs', '3',
            '--batch_size', '8',
            '--learning_rate', '1e-3',
            '--eval_interval', '1',
            '--output_dir', './checkpoints/benchmark_run'
        ]
    
    main()



