#!/bin/bash
# Vast.ai Provisioning Script — SDXL LoRA Training
# HOSTED: https://data.mochiplan.dev/vastai_provision.sh
# 
# Usage: Set env vars before instance launch in Vast.ai template:
#   DATASET_URL=https://data.mochiplan.dev/my_dataset.zip
#   DATASET_SHA256=abc123...
#   HF_MODEL_REPO=LyliaEngine/waiIllustriousSDXL_v170
#   HF_MODEL_FILE=waiIllustriousSDXL_v170.safetensors
#   MODEL_HASH=f116b0c78ff...
#   MIRROR=https://data.mochiplan.dev  # nearest mirror for tokenizer
#   TRAINING_REPO=https://github.com/kohya-ss/sd-scripts.git
#
# After this script: model at /root/model_diffusers, dataset at /root/dataset_joint/
# Ready to run sdxl_train_network.py

set -euo pipefail

echo "=== Vast.ai Provisioning Start ==="
echo "Dataset: ${DATASET_URL:-NOT SET}"
echo "Model: ${HF_MODEL_REPO:-NOT SET}/${HF_MODEL_FILE:-NOT SET}"

# ── 1. System dependencies ──
echo "[1/5] Installing system packages..."
apt-get update -qq
apt-get install -y -qq libgl1-mesa-glx libglib2.0-0 unzip wget >/dev/null 2>&1

# ── 2. Python dependencies ──
echo "[2/5] Installing Python packages..."
pip install -q diffusers accelerate transformers huggingface_hub

# ── 3. Model download & conversion ──
echo "[3/5] Downloading model from HuggingFace..."
python3 << PYEOF
from huggingface_hub import hf_hub_download
from diffusers import StableDiffusionXLPipeline
import torch, shutil, urllib.request, os

repo = os.environ.get("HF_MODEL_REPO", "LyliaEngine/waiIllustriousSDXL_v170")
file = os.environ.get("HF_MODEL_FILE", "waiIllustriousSDXL_v170.safetensors")
expected_hash = os.environ.get("MODEL_HASH", "")
mirror = os.environ.get("MIRROR", "https://data.mochiplan.dev")

path = hf_hub_download(repo, file)
if expected_hash:
    import hashlib
    actual = hashlib.sha256(open(path, "rb").read()).hexdigest()
    assert actual == expected_hash, f"MODEL HASH MISMATCH: {actual} != {expected_hash}"
    print(f"Model hash OK: {actual[:16]}...")

pipe = StableDiffusionXLPipeline.from_single_file(path, torch_dtype=torch.bfloat16)
pipe.save_pretrained("/root/model_diffusers", safe_serialization=True)

# Fix CLIP tokenizer
for src, dst in [("tokenizer_vocab.json", "vocab.json"), ("tokenizer_merges.txt", "merges.txt")]:
    url = f"{mirror}/{src}"
    for sub in ["tokenizer", "tokenizer_2"]:
        dest = f"/root/model_diffusers/{sub}/{dst}"
        urllib.request.urlretrieve(url, dest)

print("Model ready at /root/model_diffusers/")
PYEOF

# ── 4. Dataset download ──
if [ -n "${DATASET_URL:-}" ]; then
    echo "[4/5] Downloading dataset..."
    wget -q --show-progress "${DATASET_URL}" -O /root/dataset.zip
    if [ -n "${DATASET_SHA256:-}" ]; then
        echo "${DATASET_SHA256}  /root/dataset.zip" | sha256sum -c
    fi
    mkdir -p /root/dataset_joint
    unzip -q -o /root/dataset.zip -d /root/dataset_joint
    rm /root/dataset.zip
    echo "Dataset ready at /root/dataset_joint/"
else
    echo "[4/5] No DATASET_URL set — skipping."
fi

# ── 5. sd-scripts ──
echo "[5/5] Cloning sd-scripts..."
cd /root
if [ -d sd-scripts ]; then
    cd sd-scripts && git pull -q
else
    git clone -q "${TRAINING_REPO:-https://github.com/kohya-ss/sd-scripts.git}"
fi
cd /root/sd-scripts && pip install -q -r requirements.txt

echo "=== Provisioning Complete ==="
echo "Model:  /root/model_diffusers/"
echo "Dataset: /root/dataset_joint/"
echo "sd-scripts: /root/sd-scripts/"
echo "Ready to train!"
