#!/bin/bash
#SBATCH --job-name=test_encode_multi
#SBATCH --nodes=1                   # Single node
#SBATCH --ntasks-per-node=2         # Test with 2 GPUs
#SBATCH --gres=gpu:2                # Request 2 GPUs
#SBATCH --cpus-per-task=8           # CPUs per GPU
#SBATCH --mem=64G                   # Memory
#SBATCH --time=00:30:00             # 30 minutes max
#SBATCH --output=logs/test_multi_gpu_%j.out
#SBATCH --error=logs/test_multi_gpu_%j.err
# #SBATCH --partition=gpu           # Uncomment if your cluster requires a partition

# Test script for multi-GPU encoding on a single node
# This tests the distributed encoding with 100 samples across 2 GPUs

set -e

echo "=========================================="
echo "Testing Semantic Encoding (Multi-GPU)"
echo "=========================================="
echo "Job ID: $SLURM_JOB_ID"
echo "Nodes: $SLURM_JOB_NUM_NODES"
echo "Tasks: $SLURM_NTASKS"
echo "=========================================="
echo ""

# Create logs directory
mkdir -p logs

# Environment setup
echo "Setting up environment..."
source ~/.bashrc

# Activate conda environment (adjust as needed)
# conda activate your_env_name

# Set environment variables for distributed training
export MASTER_ADDR=$(scontrol show hostname $SLURM_NODELIST | head -n 1)
export MASTER_PORT=29500

echo "Master address: $MASTER_ADDR"
echo "Master port: $MASTER_PORT"
echo ""

# Configuration
FULL_METADATA="/home/tony/Data/Preference/RealGen/metas_v5_val_sampled_diverse.jsonl"
TEST_METADATA="/home/tony/Work/tony/RealGen/test_metadata_100samples.jsonl"
TEST_OUTPUT_DIR="/app2/suno/data/semantic_code/sft"
BATCH_SIZE=8


# Create test metadata with 100 samples
echo "Creating test metadata with 100 samples..."
head -n 100 $FULL_METADATA > $TEST_METADATA
echo "Test metadata created: $TEST_METADATA"
echo ""

echo "Configuration:"
echo "  Metadata: $TEST_METADATA"
echo "  Output:   $TEST_OUTPUT_DIR"
echo "  Batch:    $BATCH_SIZE"
echo "  GPUs:     2"
echo "  Samples:  100 (50 per GPU)"
echo ""

# Create output directory
mkdir -p $TEST_OUTPUT_DIR

# Launch distributed encoding
echo "Launching distributed encoding..."
echo ""

srun python /home/tony/Work/tony/RealGen/encode_semantic_codes.py \
    --metadata_path $TEST_METADATA \
    --output_dir $TEST_OUTPUT_DIR \
    --batch_size $BATCH_SIZE

echo ""
echo "=========================================="
echo "Test completed!"
echo "=========================================="
echo ""

# Check results
NUM_NPZ=$(find $TEST_OUTPUT_DIR -name "*.npz" 2>/dev/null | wc -l)
echo "Results:"
echo "  Created $NUM_NPZ .npz files"
echo "  Expected: ~100 files (or fewer if some failed)"
echo ""

if [ $NUM_NPZ -gt 0 ]; then
    echo "Sample output files:"
    ls -lh $TEST_OUTPUT_DIR/*.npz | head -n 3
    echo ""
    echo "✓ Multi-GPU test passed!"
else
    echo "✗ No output files created. Check error logs."
    exit 1
fi

