File size: 2,153 Bytes
3de85f5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 | #!/bin/bash
#SBATCH --job-name=nep_HfO2_16card
#SBATCH --partition=hx1hdexclu12
#SBATCH --nodes=2
#SBATCH --ntasks-per-node=8
#SBATCH --gres=dcu:8
#SBATCH --cpus-per-task=4
#SBATCH --mem=256G
#SBATCH --time=2:00:00
#SBATCH --output=slurm_16card_%j.out
#SBATCH --error=slurm_16card_%j.err
if [[ -n "${SCRIPT_DIR}" ]]; then
:
elif [[ -n "${SLURM_SUBMIT_DIR}" ]]; then
SCRIPT_DIR="${SLURM_SUBMIT_DIR}"
else
echo "ERROR: 未检测到外部 SCRIPT_DIR 变量,也不在 Slurm 任务环境(SLURM_SUBMIT_DIR 为空)"
exit 1
fi
echo $SCRIPT_DIR
export MATCHEM_CONDA_NAME="${MATCHEM_CONDA_NAME:-test_pip}"
source "$SCRIPT_DIR/matchem_env.sh"
# MatPL 运行时环境(不侵入 matchem_env.sh,各算例自行维护)
# 校验MatPL根路径
if [[ -z "${MATPL_SRC_DIR}" || "${MATPL_SRC_DIR}" == "/path/to/matpl_dcu" || ! -d "${MATPL_SRC_DIR}" ]]; then
echo "=============================================="
echo "ERROR: MATPL_SRC_DIR 路径配置错误!"
echo "当前值:${MATPL_SRC_DIR}"
echo "请设置为真实有效的MatPL源码目录,示例:"
echo "export MATPL_SRC_DIR=/public/home/xxx/real_matpl_path"
echo "=============================================="
exit 1
fi
if [ -f "$MATPL_SRC_DIR/env.sh" ]; then
source "$MATPL_SRC_DIR/env.sh"
fi
echo "MATPL_SRC_DIR = $MATPL_SRC_DIR"
export LD_LIBRARY_PATH="$MATPL_SRC_DIR/src/op/build/lib:${LD_LIBRARY_PATH:-}"
# 多节点通信设置
MASTER_ADDR=$(scontrol show hostnames $SLURM_JOB_NODELIST | head -n 1)
function get_free_port() {
python -c 'import socket; s = socket.socket(socket.AF_INET, socket.SOCK_STREAM); s.bind(("", 0)); print(s.getsockname()[1]); s.close()'
}
MASTER_PORT=$(get_free_port)
export MASTER_ADDR=$MASTER_ADDR
export MASTER_PORT=$MASTER_PORT
# 16 卡训练(2 节点 × 8 DCU)
# 输入 JSON 使用环境变量,运行前展开为临时文件
INPUT_JSON="HfO2_nep_train_16card.json"
EXPANDED_JSON=".${INPUT_JSON%.json}_expanded_$$.json"
trap 'rm -f "$EXPANDED_JSON"' EXIT
python3 -c "import os; open('$EXPANDED_JSON','w').write(os.path.expandvars(open('$INPUT_JSON').read()))"
srun MatPL train "$EXPANDED_JSON"
|