#!/bin/bash
# ============================================
# MindSpeed-Bridge Docker Image Build Script
# ============================================

set -e

NPU_TYPE="910b"
IMAGE_NAME=""
OS="openeuler24.03"
BASE_IMAGE=""
PYTHON_VERSION="3.12"
TORCH_VERSION="2.9.0"
TORCH_NPU_VERSION="2.9.0"
TORCH_WHL_URL=""
TORCH_NPU_WHL_URL="https://gitcode.com/Ascend/pytorch/releases/download/v26.1.0-beta.2-pytorch2.9.0/torch_npu-2.9.0.post5-cp312-cp312-manylinux_2_28_aarch64.whl"
TRITON_ASCEND_VERSION="3.2.1"
BASE_IMAGE_VERSION="9.1.0-beta.3"
MINDSPEED_BRIDGE_BRANCH="master"
MINDSPEED_BRANCH="core_r0.16.0"
MINDSPEED_OPS_BRANCH="master"
MEGATRON_BRANCH="core_v0.16.1"
MEGATRON_BRIDGE_BRANCH="v0.3.1"
FLASH_LINEAR_ATTENTION_NPU_BRANCH="v26.1.0"
FLA_NPU_SOC=""
FLA_NPU_OPS=(
    causal_conv1d
    chunk_bwd_dv_local
    chunk_bwd_dqkwg
    chunk_gated_delta_rule_bwd_dhu
    prepare_wy_repr_bwd_da
    prepare_wy_repr_bwd_full
    chunk_fwd_o
    chunk_gated_delta_rule_fwd_h
    recurrent_gated_delta_rule
    recompute_wu_fwd
)
NO_CACHE=""
NPU_TYPE_EXPLICIT=false
OS_EXPLICIT=false
HTTP_PROXY_VALUE="${HTTP_PROXY:-${http_proxy:-}}"
HTTPS_PROXY_VALUE="${HTTPS_PROXY:-${https_proxy:-}}"
NO_PROXY_VALUE="${NO_PROXY:-${no_proxy:-}}"

show_help() {
    cat << EOF
Usage: $0 [OPTIONS]

Build MindSpeed-Bridge Docker image.

Options:
    -t, --npu-type TYPE              NPU type: 910b, a3 or 950 (default: 910b)
    -i, --image-name NAME            Custom output image full name
    -o, --os OS                      OS: openeuler24.03 or ubuntu22.04
    -n, --no-cache                   Build without Docker cache
    --base-image IMAGE               Full base image name. If set, passed to FROM as-is
    --base-image-version VER         AscendHub CANN base image version (default: 9.1.0-beta.3)
    --python-version VER             Python version (default: 3.12)
    --torch-version VER              PyTorch version (default: 2.9.0)
    --torch-npu-version VER          TorchNPU version (default: 2.9.0)
    --torch-whl-url URL              Install torch from a specific wheel URL
    --torch-npu-whl-url URL          Install torch_npu from a specific wheel URL
    --triton-ascend-version VER      triton-ascend version (default: 3.2.1)
    --mindspeed-bridge-branch VER    MindSpeed-Bridge git branch/version (default: master)
    --mindspeed-branch VER           MindSpeed git branch/version (default: core_r0.16.0)
    --mindspeed-ops-branch VER       MindSpeed-Ops git branch/version (default: master)
    --megatron-branch VER            Megatron-LM git branch/version (default: core_v0.16.1)
    --megatron-bridge-branch VER     Megatron-Bridge git branch/version (default: v0.3.1)
    --fla-npu-branch VER             flash-linear-attention-npu git branch/version (default: v26.1.0)
    --fla-npu-soc SOC                flash-linear-attention-npu build soc. Default is mapped from NPU type:
                                     910b -> ascend910b, a3 -> ascend910_93, 950 -> ascend950
    -h, --help                       Show this help message and exit

Image tag convention:
    {mindspeed_bridge_branch}-{npu_type}-{os}-cann{base_image_version}-torchnpu{torch_npu_version}-py{python_version}-{arch}

Examples:
    bash image_build.sh
    bash image_build.sh -t 910b --base-image-version 9.1.0-beta.3
    bash image_build.sh --base-image swr.cn-south-1.myhuaweicloud.com/ascendhub/cann:9.1.0-beta.3-910b-openeuler24.03-py3.12
    bash image_build.sh --torch-npu-whl-url https://gitcode.com/Ascend/pytorch/releases/download/v26.1.0-beta.2-pytorch2.9.0/torch_npu-2.9.0.post5-cp312-cp312-manylinux_2_28_aarch64.whl
EOF
}

parse_base_image_tag() {
    local image="$1"
    local tag="${image##*:}"
    local tag_lower
    tag_lower=$(echo "$tag" | tr '[:upper:]' '[:lower:]')

    if [[ "$tag_lower" == *"910b"* ]]; then
        DETECTED_NPU_TYPE="910b"
    elif [[ "$tag_lower" == *"-950-"* ]]; then
        DETECTED_NPU_TYPE="950"
    elif [[ "$tag_lower" == *"-a3-"* ]] || [[ "$tag_lower" == *"-a3-py"* ]]; then
        DETECTED_NPU_TYPE="a3"
    fi

    if [[ "$tag_lower" == *"openeuler24.03"* ]]; then
        DETECTED_OS="openeuler24.03"
    elif [[ "$tag_lower" == *"ubuntu22.04"* ]]; then
        DETECTED_OS="ubuntu22.04"
    fi

    if [[ "$tag_lower" =~ py([0-9]+\.[0-9]+) ]]; then
        DETECTED_PYTHON_VERSION="${BASH_REMATCH[1]}"
    fi
}

while [[ $# -gt 0 ]]; do
    case $1 in
        -t|--npu-type)              NPU_TYPE="$2"; NPU_TYPE_EXPLICIT=true; shift 2 ;;
        -i|--image-name)            IMAGE_NAME="$2"; shift 2 ;;
        -o|--os)                    OS="$2"; OS_EXPLICIT=true; shift 2 ;;
        -n|--no-cache)              NO_CACHE="--no-cache"; shift ;;
        --base-image)               BASE_IMAGE="$2"; shift 2 ;;
        --base-image-version)       BASE_IMAGE_VERSION="$2"; shift 2 ;;
        --python-version)           PYTHON_VERSION="$2"; shift 2 ;;
        --torch-version)            TORCH_VERSION="$2"; shift 2 ;;
        --torch-npu-version)        TORCH_NPU_VERSION="$2"; shift 2 ;;
        --torch-whl-url)            TORCH_WHL_URL="$2"; shift 2 ;;
        --torch-npu-whl-url)        TORCH_NPU_WHL_URL="$2"; shift 2 ;;
        --triton-ascend-version)    TRITON_ASCEND_VERSION="$2"; shift 2 ;;
        --mindspeed-bridge-branch)  MINDSPEED_BRIDGE_BRANCH="$2"; shift 2 ;;
        --mindspeed-branch)         MINDSPEED_BRANCH="$2"; shift 2 ;;
        --mindspeed-ops-branch)     MINDSPEED_OPS_BRANCH="$2"; shift 2 ;;
        --megatron-branch)          MEGATRON_BRANCH="$2"; shift 2 ;;
        --megatron-bridge-branch)   MEGATRON_BRIDGE_BRANCH="$2"; shift 2 ;;
        --fla-npu-branch)           FLASH_LINEAR_ATTENTION_NPU_BRANCH="$2"; shift 2 ;;
        --fla-npu-soc)              FLA_NPU_SOC="$2"; shift 2 ;;
        -h|--help)                  show_help; exit 0 ;;
        *)                          echo "Unknown argument: $1"; show_help; exit 1 ;;
    esac
done

SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
DOCKERFILE="${SCRIPT_DIR}/Dockerfile"

if [ ! -f "$DOCKERFILE" ]; then
    echo "Error: Dockerfile not found: $DOCKERFILE"
    exit 1
fi

DETECTED_NPU_TYPE=""
DETECTED_OS=""
DETECTED_PYTHON_VERSION=""
if [ -n "$BASE_IMAGE" ]; then
    parse_base_image_tag "$BASE_IMAGE"
    if [ "$NPU_TYPE_EXPLICIT" = false ] && [ -n "$DETECTED_NPU_TYPE" ]; then
        NPU_TYPE="$DETECTED_NPU_TYPE"
    fi
    if [ "$OS_EXPLICIT" = false ] && [ -n "$DETECTED_OS" ]; then
        OS="$DETECTED_OS"
    fi
    if [ -n "$DETECTED_PYTHON_VERSION" ]; then
        PYTHON_VERSION="$DETECTED_PYTHON_VERSION"
    fi
fi

NPU_TYPE_LOWER=$(echo "$NPU_TYPE" | tr '[:upper:]' '[:lower:]')
OS=$(echo "$OS" | tr '[:upper:]' '[:lower:]')

if [ "$NPU_TYPE_LOWER" != "910b" ] && [ "$NPU_TYPE_LOWER" != "a3" ] && [ "$NPU_TYPE_LOWER" != "950" ]; then
    echo "Error: NPU type must be 910b, a3 or 950"
    exit 1
fi

if [ -z "$FLA_NPU_SOC" ]; then
    case "$NPU_TYPE_LOWER" in
        910b) FLA_NPU_SOC="ascend910b" ;;
        a3)   FLA_NPU_SOC="ascend910_93" ;;
        950)  FLA_NPU_SOC="ascend950" ;;
    esac
fi

if [ "$OS" != "ubuntu22.04" ] && [ "$OS" != "openeuler24.03" ]; then
    echo "Error: OS must be ubuntu22.04 or openeuler24.03"
    exit 1
fi

case "$OS" in
    ubuntu*) OS_FAMILY="ubuntu"; REPO_SCRIPT="configure_apt_repo.sh" ;;
    openeuler*) OS_FAMILY="openeuler"; REPO_SCRIPT="configure_yum_repo.sh" ;;
esac

HOST_ARCH=$(uname -m)
case "$HOST_ARCH" in
    arm64) ARCH_NAME="aarch64" ;;
    *)     ARCH_NAME="$HOST_ARCH" ;;
esac

if [ -z "$IMAGE_NAME" ]; then
    TAG_REF=$(echo "$MINDSPEED_BRIDGE_BRANCH" | tr '/:' '--')
    TORCH_NPU_VERSION_TAG=$(echo "$TORCH_NPU_VERSION" | sed 's/^v//' | tr '/:' '--')
    IMAGE_NAME="mindspeed-bridge:${TAG_REF}-${NPU_TYPE_LOWER}-${OS}-cann${BASE_IMAGE_VERSION}-torchnpu${TORCH_NPU_VERSION_TAG}-py${PYTHON_VERSION}-${ARCH_NAME}"
fi

cd "$SCRIPT_DIR"
cp "${SCRIPT_DIR}/${REPO_SCRIPT}" configure_repo.sh
trap 'rm -f configure_repo.sh' EXIT

BUILD_ARGS="--build-arg OS=${OS}"
BUILD_ARGS="$BUILD_ARGS --build-arg OS_FAMILY=${OS_FAMILY}"
BUILD_ARGS="$BUILD_ARGS --build-arg NPU_TYPE=${NPU_TYPE_LOWER}"
BUILD_ARGS="$BUILD_ARGS --build-arg PYTHON_VERSION=${PYTHON_VERSION}"
BUILD_ARGS="$BUILD_ARGS --build-arg TORCH_VERSION=${TORCH_VERSION}"
BUILD_ARGS="$BUILD_ARGS --build-arg TORCH_NPU_VERSION=${TORCH_NPU_VERSION}"
BUILD_ARGS="$BUILD_ARGS --build-arg TORCH_WHL_URL=${TORCH_WHL_URL}"
BUILD_ARGS="$BUILD_ARGS --build-arg TORCH_NPU_WHL_URL=${TORCH_NPU_WHL_URL}"
BUILD_ARGS="$BUILD_ARGS --build-arg TRITON_ASCEND_VERSION=${TRITON_ASCEND_VERSION}"
BUILD_ARGS="$BUILD_ARGS --build-arg MINDSPEED_BRIDGE_BRANCH=${MINDSPEED_BRIDGE_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg MINDSPEED_BRANCH=${MINDSPEED_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg MINDSPEED_OPS_BRANCH=${MINDSPEED_OPS_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg MEGATRON_BRANCH=${MEGATRON_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg MEGATRON_BRIDGE_BRANCH=${MEGATRON_BRIDGE_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg FLASH_LINEAR_ATTENTION_NPU_BRANCH=${FLASH_LINEAR_ATTENTION_NPU_BRANCH}"
BUILD_ARGS="$BUILD_ARGS --build-arg FLA_NPU_SOC=${FLA_NPU_SOC}"
FLA_NPU_OPS_CSV=$(IFS=,; echo "${FLA_NPU_OPS[*]}")
BUILD_ARGS="$BUILD_ARGS --build-arg FLA_NPU_OPS=${FLA_NPU_OPS_CSV}"

PROXY_BUILD_ARGS=()
if [ -n "$HTTP_PROXY_VALUE" ]; then
    PROXY_BUILD_ARGS+=(--build-arg "http_proxy=${HTTP_PROXY_VALUE}")
    PROXY_BUILD_ARGS+=(--build-arg "HTTP_PROXY=${HTTP_PROXY_VALUE}")
fi
if [ -n "$HTTPS_PROXY_VALUE" ]; then
    PROXY_BUILD_ARGS+=(--build-arg "https_proxy=${HTTPS_PROXY_VALUE}")
    PROXY_BUILD_ARGS+=(--build-arg "HTTPS_PROXY=${HTTPS_PROXY_VALUE}")
fi
if [ -n "$NO_PROXY_VALUE" ]; then
    PROXY_BUILD_ARGS+=(--build-arg "no_proxy=${NO_PROXY_VALUE}")
    PROXY_BUILD_ARGS+=(--build-arg "NO_PROXY=${NO_PROXY_VALUE}")
fi

if [ -n "$BASE_IMAGE" ]; then
    BUILD_ARGS="$BUILD_ARGS --build-arg BASE_IMAGE=${BASE_IMAGE}"
else
    BUILD_ARGS="$BUILD_ARGS --build-arg BASE_IMAGE_VERSION=${BASE_IMAGE_VERSION}"
fi

echo "=========================================="
echo "Build Configuration"
echo "=========================================="
echo "NPU Type:              ${NPU_TYPE_LOWER}"
echo "Image Name:            ${IMAGE_NAME}"
echo "OS:                    ${OS}"
echo "OS_FAMILY:             ${OS_FAMILY}"
echo "Base Image Version:    ${BASE_IMAGE_VERSION}"
if [ -n "$BASE_IMAGE" ]; then
    echo "Base Image:            ${BASE_IMAGE}"
fi
echo "Python Version:        ${PYTHON_VERSION}"
echo "PyTorch Version:       ${TORCH_VERSION}"
echo "TorchNPU Version:      ${TORCH_NPU_VERSION}"
echo "triton-ascend Ver:     ${TRITON_ASCEND_VERSION}"
echo "MindSpeed-Bridge Ver:  ${MINDSPEED_BRIDGE_BRANCH}"
echo "MindSpeed Ver:         ${MINDSPEED_BRANCH}"
echo "MindSpeed-Ops Ver:     ${MINDSPEED_OPS_BRANCH}"
echo "Megatron-LM Ver:       ${MEGATRON_BRANCH}"
echo "Megatron-Bridge Ver:   ${MEGATRON_BRIDGE_BRANCH}"
echo "FLA NPU Ver:           ${FLASH_LINEAR_ATTENTION_NPU_BRANCH}"
echo "FLA NPU SOC:           ${FLA_NPU_SOC}"
echo "FLA NPU Ops:           ${FLA_NPU_OPS_CSV}"
if [ -n "$HTTP_PROXY_VALUE" ] || [ -n "$HTTPS_PROXY_VALUE" ]; then
    echo "Proxy:                 enabled"
else
    echo "Proxy:                 disabled"
fi
echo "No Cache:              ${NO_CACHE:-No}"
echo "=========================================="

docker build \
    -t "$IMAGE_NAME" \
    -f "$DOCKERFILE" \
    $BUILD_ARGS \
    "${PROXY_BUILD_ARGS[@]}" \
    $NO_CACHE \
    --network=host \
    .

echo ""
echo "=========================================="
echo "Build Complete!"
echo "Image: ${IMAGE_NAME}"
echo "=========================================="