#!/bin/bash

# Setup Suno Environment with Flash Attention v3 (Hopper) and v2 (fixed commit)
# This version installs FA3 first (for H100 optimization) then FA2 for compatibility
# FA3 is compatible with FA2, so both can coexist

set -e  # Exit on error

# Configuration
ENV_NAME="${1:-suno_env_fa2_fa3}"
SUNO_UTILS_PATH="${2}"
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
LOG_FILE="${SCRIPT_DIR}/tmp/setup_${ENV_NAME}_$(date +%Y%m%d_%H%M%S).log"

# Create tmp directory if it doesn't exist
mkdir -p "${SCRIPT_DIR}/tmp"

# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
BLUE='\033[0;34m'
NC='\033[0m' # No Color

# Logging functions
log() {
    echo -e "${GREEN}[$(date +'%Y-%m-%d %H:%M:%S')]${NC} $1" | tee -a "$LOG_FILE"
}

error() {
    echo -e "${RED}[ERROR]${NC} $1" | tee -a "$LOG_FILE"
    exit 1
}

warning() {
    echo -e "${YELLOW}[WARNING]${NC} $1" | tee -a "$LOG_FILE"
}

info() {
    echo -e "${BLUE}[INFO]${NC} $1" | tee -a "$LOG_FILE"
}

# Function to validate Suno Utils path
setup_suno_utils_path() {
    # Check if suno_utils path was provided
    if [ -z "$SUNO_UTILS_PATH" ]; then
        error "Suno utils path is required as second argument.\nUsage: $0 <env_name> <suno_utils_path>\nExample: $0 suno_env_auto ~/projects/glockenspiel/suno_utils"
    fi

    # Expand tilde if present
    SUNO_UTILS_PATH="${SUNO_UTILS_PATH/#\~/$HOME}"

    log "Using suno_utils path: $SUNO_UTILS_PATH"

    # Check if suno_utils exists
    if [ ! -d "$SUNO_UTILS_PATH" ]; then
        error "suno_utils not found at $SUNO_UTILS_PATH. Please provide a valid path to the suno_utils package."
    fi

    log "suno_utils found at $SUNO_UTILS_PATH"

    # Check if it's a valid Python package
    if [ ! -f "$SUNO_UTILS_PATH/setup.py" ] && [ ! -f "$SUNO_UTILS_PATH/pyproject.toml" ]; then
        error "$SUNO_UTILS_PATH doesn't appear to be a valid Python package (missing setup.py or pyproject.toml)"
    fi
}

# Function to setup Flash Attention v3 path (main repository with fixed commit)
setup_flash_attention_v3_path() {
    # Always clone fresh to /tmp with specific commit (newer commit for both FA3 and FA2)
    FLASH_ATTN_V3_COMMIT="7b0bfcc3d1f69786f0c4277c582ad58acdfb297d"
    FLASH_ATTN_V3_PATH="/tmp/flash-attention-v3-${ENV_NAME}"

    log "Cloning Flash Attention v3 to ${FLASH_ATTN_V3_PATH} (commit ${FLASH_ATTN_V3_COMMIT})..."

    # Remove if exists
    rm -rf "$FLASH_ATTN_V3_PATH" >> "$LOG_FILE" 2>&1 || true

    # Clone repository
    git clone https://github.com/Dao-AILab/flash-attention.git "$FLASH_ATTN_V3_PATH" >> "$LOG_FILE" 2>&1 || {
        error "Failed to clone Flash Attention v3 repository"
    }

    # Checkout specific commit
    cd "$FLASH_ATTN_V3_PATH"
    log "Checking out commit ${FLASH_ATTN_V3_COMMIT}..."
    git checkout "$FLASH_ATTN_V3_COMMIT" >> "$LOG_FILE" 2>&1 || {
        error "Failed to checkout commit ${FLASH_ATTN_V3_COMMIT}"
    }

    # Initialize submodules
    log "Initializing Flash Attention v3 submodules..."
    git submodule update --init --recursive >> "$LOG_FILE" 2>&1 || {
        warning "Failed to initialize some submodules, continuing anyway"
    }
    cd - > /dev/null

    log "Flash Attention v3 cloned successfully at commit ${FLASH_ATTN_V3_COMMIT}!"

    # Check if hopper directory exists (for FA v3)
    if [ ! -d "$FLASH_ATTN_V3_PATH/hopper" ]; then
        warning "Flash Attention v3 hopper directory not found at $FLASH_ATTN_V3_PATH/hopper"
        warning "This commit might not have the hopper subdirectory for H100 optimization"
    fi
}

# Function to setup Flash Attention v2 path (fixed commit)
setup_flash_attention_v2_path() {
    # Always clone fresh to /tmp with specific commit
    FLASH_ATTN_COMMIT="7b0bfcc3d1f69786f0c4277c582ad58acdfb297d"
    FLASH_ATTN_PATH="/tmp/flash-attention-${ENV_NAME}"

    log "Cloning Flash Attention v2 to ${FLASH_ATTN_PATH} (commit ${FLASH_ATTN_COMMIT})..."

    # Remove if exists
    rm -rf "$FLASH_ATTN_PATH" >> "$LOG_FILE" 2>&1 || true

    # Clone repository
    git clone https://github.com/Dao-AILab/flash-attention.git "$FLASH_ATTN_PATH" >> "$LOG_FILE" 2>&1 || {
        error "Failed to clone Flash Attention v2 repository"
    }

    # Checkout specific commit
    cd "$FLASH_ATTN_PATH"
    log "Checking out commit ${FLASH_ATTN_COMMIT}..."
    git checkout "$FLASH_ATTN_COMMIT" >> "$LOG_FILE" 2>&1 || {
        error "Failed to checkout commit ${FLASH_ATTN_COMMIT}"
    }

    # Initialize submodules
    log "Initializing Flash Attention v2 submodules..."
    git submodule update --init --recursive >> "$LOG_FILE" 2>&1 || {
        warning "Failed to initialize some submodules, continuing anyway"
    }
    cd - > /dev/null

    log "Flash Attention v2 cloned successfully at commit ${FLASH_ATTN_COMMIT}!"
}

# Main setup function
main() {
    log "Starting Suno Environment Setup with Flash Attention v3 + v2"
    log "Environment name: ${ENV_NAME}"
    log "Log file: ${LOG_FILE}"
    log "Strategy: Install FA3 (H100 optimized) first, then FA2 (fixed commit) for compatibility"

    # Check CUDA version
    log "Checking CUDA version..."
    if command -v nvcc &> /dev/null; then
        cuda_version=$(nvcc --version | grep "release" | sed 's/.*release //' | cut -d',' -f1)
        log "CUDA version: ${cuda_version}"

        # Flash Attention v3 requires CUDA >= 12.3
        cuda_major=$(echo $cuda_version | cut -d'.' -f1)
        cuda_minor=$(echo $cuda_version | cut -d'.' -f2)

        if [ "$cuda_major" -lt 12 ] || ([ "$cuda_major" -eq 12 ] && [ "$cuda_minor" -lt 3 ]); then
            warning "Flash Attention v3 requires CUDA >= 12.3, found ${cuda_version}"
            warning "Installation may fail or performance may be suboptimal"
        fi
    else
        warning "nvcc not found. Unable to check CUDA version."
        warning "Flash Attention v3 requires CUDA >= 12.3"
    fi

    # Setup paths
    setup_suno_utils_path
    setup_flash_attention_v3_path
    setup_flash_attention_v2_path

    # Check if environment already exists
    if conda env list | grep -q "^${ENV_NAME} "; then
        warning "Environment ${ENV_NAME} already exists."
        read -p "Do you want to remove and recreate it? (y/n): " -n 1 -r
        echo
        if [[ $REPLY =~ ^[Yy]$ ]]; then
            log "Removing existing environment..."
            conda env remove -n "${ENV_NAME}" -y >> "$LOG_FILE" 2>&1
        else
            error "Environment already exists. Exiting."
        fi
    fi

    # Step 1: Create conda environment
    log "Step 1: Creating conda environment with Python 3.10..."
    conda create -n "${ENV_NAME}" python=3.10.15 -y >> "$LOG_FILE" 2>&1
    log "Conda environment created!"

    # Activate the environment
    source "$(conda info --base)/etc/profile.d/conda.sh"
    conda activate "${ENV_NAME}"

    # Step 2: Install suno_utils first
    log "Step 2: Installing suno_utils from ${SUNO_UTILS_PATH}..."
    pip install -e "$SUNO_UTILS_PATH" >> "$LOG_FILE" 2>&1
    log "suno_utils installed!"

    # Step 3: Check which torch version was installed
    log "Step 3: Checking torch version installed by suno_utils..."
    EXISTING_TORCH=$(python -c "import torch; print(torch.__version__)" 2>/dev/null || echo "None")
    log "Torch version installed by suno_utils: ${EXISTING_TORCH}"

    # Step 4: Uninstall torch (whatever version suno_utils installed)
    log "Step 4: Uninstalling torch installed by suno_utils..."
    pip uninstall torch torchvision torchaudio -y >> "$LOG_FILE" 2>&1
    log "Torch uninstalled!"

    # Step 5: Install PyTorch 2.6.0 with CUDA 12.4 support
    log "Step 5: Installing PyTorch 2.6.0 with CUDA 12.4 support..."
    pip install torch==2.6.0 torchvision==0.21.0 torchaudio==2.6.0 --index-url https://download.pytorch.org/whl/cu124 >> "$LOG_FILE" 2>&1
    log "PyTorch installed with CUDA support!"

    # Verify PyTorch installation
    python -c "import torch; print(f'PyTorch version: {torch.__version__}')"
    python -c "import torch; print(f'CUDA available: {torch.cuda.is_available()}')"

    # Step 6: Build Flash Attention v3 from hopper directory
    log "Step 6: Building Flash Attention v3 (H100 optimized) from ${FLASH_ATTN_V3_PATH}/hopper..."

    if [ ! -d "$FLASH_ATTN_V3_PATH/hopper" ]; then
        warning "Flash Attention v3 hopper directory not found at $FLASH_ATTN_V3_PATH/hopper"
        warning "Skipping FA3 installation, will only install FA2"
    else
        cd "$FLASH_ATTN_V3_PATH/hopper"

        # Clean any previous builds
        log "Cleaning previous Flash Attention v3 builds..."
        rm -rf build dist *.egg-info >> "$LOG_FILE" 2>&1 || true

        # Install required packages
        log "Installing build dependencies..."
        pip install packaging ninja >> "$LOG_FILE" 2>&1

        # Install Flash Attention v3
        log "Building and installing Flash Attention v3 (this may take 10-30 minutes)..."
        log "Note: Flash Attention v3 is optimized for H100/H800 GPUs"

        export MAX_JOBS=4  # Limit parallel jobs
        export FLASH_ATTENTION_FORCE_BUILD=TRUE

        # Use python setup.py install as recommended in README
        python setup.py install >> "$LOG_FILE" 2>&1 || {
            warning "Flash Attention v3 installation failed"
            warning "This is expected if not running on H100/H800"
            warning "Continuing with FA2 installation..."
        }

        # Test Flash Attention v3 import
        log "Testing Flash Attention v3 installation..."
        python -c "
try:
    import flash_attn_interface
    print('✓ Flash Attention v3 successfully installed')
    print(f'  Module location: {flash_attn_interface.__file__}')
except ImportError as e:
    print('✗ Flash Attention v3 not available')
    print(f'  Error: {e}')
" | tee -a "$LOG_FILE"
    fi

    # Step 7: Build Flash Attention v2 from fixed commit
    log "Step 7: Building Flash Attention v2 from ${FLASH_ATTN_PATH} (commit ${FLASH_ATTN_COMMIT})..."

    cd "$FLASH_ATTN_PATH"

    # Clean any previous builds
    log "Cleaning previous Flash Attention v2 builds..."
    rm -rf build dist *.egg-info >> "$LOG_FILE" 2>&1 || true

    # Install ninja if not already installed
    log "Ensuring ninja is installed..."
    pip install ninja >> "$LOG_FILE" 2>&1

    # Install Flash Attention v2
    log "Building and installing Flash Attention v2 (this may take 10-30 minutes)..."
    export MAX_JOBS=4  # Limit parallel jobs
    export FLASH_ATTENTION_FORCE_BUILD=TRUE
    pip install . --no-build-isolation >> "$LOG_FILE" 2>&1 || {
        error "Flash Attention v2 installation failed"
    }

    # Test Flash Attention v2 import
    log "Testing Flash Attention v2 installation..."
    python -c "
try:
    import flash_attn
    print('✓ Flash Attention v2 successfully installed')
    print(f'  Version: {flash_attn.__version__}')
    print(f'  Module location: {flash_attn.__file__}')
except ImportError as e:
    print('✗ Flash Attention v2 not available')
    print(f'  Error: {e}')
" | tee -a "$LOG_FILE"

    # Step 8: Install additional Python packages with transformers 4.57.0
    log "Step 8: Installing additional Python packages..."
    pip install wandb nnAudio deepspeed auraloss torchsde g2p_en \
        transformers==4.57.0 modal==1.0.1 better_profanity encodec \
        pytorch-lightning sentencepiece tiktoken einops ffmpeg-python >> "$LOG_FILE" 2>&1
    log "Additional packages installed!"

    # Step 9: Install sox via conda
    log "Step 9: Installing sox via conda..."
    conda install -c conda-forge sox -y >> "$LOG_FILE" 2>&1
    log "Sox installed!"

    # Return to script directory
    cd "${SCRIPT_DIR}"

    # Verification
    log "Running verification tests..."
    if [ -f "${SCRIPT_DIR}/verify_imports.py" ]; then
        python "${SCRIPT_DIR}/verify_imports.py" | tee -a "$LOG_FILE"
    else
        warning "verify_imports.py not found, skipping verification"
    fi

    # Create activation script
    log "Creating activation script..."
    cat > "${SCRIPT_DIR}/activate_${ENV_NAME}.sh" << EOF
#!/bin/bash
# Activation script for ${ENV_NAME} environment with Flash Attention v3 + v2

# Source conda
source "\$(conda info --base)/etc/profile.d/conda.sh"

# Activate environment
conda activate ${ENV_NAME}

# Set environment variables
export CUDA_VISIBLE_DEVICES=\${CUDA_VISIBLE_DEVICES:-0}
export OMP_NUM_THREADS=1
export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_\$USER
export PYTHONPATH="${FLASH_ATTN_V3_PATH}/hopper:\$PYTHONPATH"

echo "Environment ${ENV_NAME} activated (Flash Attention v3 + v2)!"
echo "Python: \$(which python)"
echo "PyTorch: \$(python -c 'import torch; print(torch.__version__)' 2>/dev/null || echo 'Not available')"
echo "CUDA available: \$(python -c 'import torch; print(torch.cuda.is_available())' 2>/dev/null || echo 'Unknown')"
echo "Flash Attention v3 path: ${FLASH_ATTN_V3_PATH}"
echo "Flash Attention v2 path: ${FLASH_ATTN_PATH}"
echo "Suno Utils path: ${SUNO_UTILS_PATH}"

# Test Flash Attention versions
echo ""
echo "Flash Attention Status:"
python -c "
try:
    import flash_attn_interface
    print('  ✓ FA3 (H100 optimized): Available')
except ImportError:
    print('  ✗ FA3: Not available')

try:
    import flash_attn
    print('  ✓ FA2 (commit ${FLASH_ATTN_COMMIT}): Available')
except ImportError:
    print('  ✗ FA2: Not available')
" 2>/dev/null
EOF

    chmod +x "${SCRIPT_DIR}/activate_${ENV_NAME}.sh"

    log "================================================"
    log "Environment setup completed!"
    log "Flash Attention v3 path: ${FLASH_ATTN_V3_PATH}"
    log "Flash Attention v2 path: ${FLASH_ATTN_PATH}"
    log "Flash Attention v2 commit: ${FLASH_ATTN_COMMIT}"
    log "Suno Utils path: ${SUNO_UTILS_PATH}"
    log ""
    log "To activate the environment, run:"
    log "  source ${SCRIPT_DIR}/activate_${ENV_NAME}.sh"
    log "Or:"
    log "  conda activate ${ENV_NAME}"
    log ""
    log "Note: This environment has both FA3 (H100 optimized) and FA2 (compatibility)"
    log "================================================"
}

# Run main function
main "$@"
