pi05_pour_full / configs /deployment_serve_policy.sh
giakhuyendihoc's picture
Upload UR5 full fine-tuned checkpoint: pi05_pour_full configs/deployment_serve_policy.sh
4bf51a3 verified
Raw
History Blame Contribute Delete
7.04 kB
#!/usr/bin/env bash
set -euo pipefail
ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
cd "$ROOT"
TASK="${TASK:-}"
RTC_WAS_SET_FOR_SERVER="${RTC+x}"
if [[ -n "$TASK" ]]; then
TASK_ENV="$ROOT/deployment/tasks/${TASK}.env"
LOCAL_TASK_ENV="$ROOT/local_pi_mamba/deployment_tasks/${TASK}.env"
if [[ ! -f "$TASK_ENV" ]]; then
if [[ -f "$LOCAL_TASK_ENV" ]]; then
TASK_ENV="$LOCAL_TASK_ENV"
else
echo "Unknown task preset: $TASK" >&2
echo "Expected: $TASK_ENV" >&2
echo " or: $LOCAL_TASK_ENV" >&2
exit 2
fi
fi
source "$TASK_ENV"
fi
OPENPI_REPO="${OPENPI_REPO:-$ROOT/openpi}"
OPENPI_VENV="${OPENPI_VENV:-$ROOT/.venvs/openpi}"
CHECKPOINT_DIR="${CHECKPOINT_DIR:-}"
CONFIG="${CONFIG:-}"
MODEL_NAME="${MODEL_NAME:-$(basename "$(dirname "${CHECKPOINT_DIR:-checkpoint}/x")")}"
PORT="${PORT:-8018}"
DEVICE="${DEVICE:-cuda}"
NUM_DENOISE_STEPS="${NUM_DENOISE_STEPS:-5}"
TMUX="${TMUX:-1}"
SESSION_NAME="${SESSION_NAME:-vla_ur5_serve_${TASK:-policy}_${PORT}}"
LOG_ROOT="${LOG_ROOT:-/mnt/vla_shared/vla_ur5/logs}"
CHECKPOINT_STEP="$(basename "${CHECKPOINT_DIR:-0}")"
METRICS_TASK_NAME="${TASK_NAME:-${TASK:-policy}}"
REFERENCE_METRICS_JSON="${REFERENCE_METRICS_JSON:-$ROOT/reports/$METRICS_TASK_NAME/metrics/${MODEL_NAME}_${CHECKPOINT_STEP}_metrics.json}"
if [[ -n "$RTC_WAS_SET_FOR_SERVER" ]]; then
echo "Note: RTC=${RTC:-unset} was passed to serve_policy.sh, but RTC is a robot-client flag." >&2
echo " Set RTC on deployment/scripts/run_robot_trial.sh or run_robot_policy.sh; the server can stay running." >&2
fi
if [[ -z "$CHECKPOINT_DIR" || -z "$CONFIG" ]]; then
echo "Set CHECKPOINT_DIR=<checkpoint-step-dir> and CONFIG=<openpi-config-name>, or set TASK=<preset>." >&2
exit 2
fi
if [[ ! -d "$CHECKPOINT_DIR" ]]; then
echo "Checkpoint directory does not exist: $CHECKPOINT_DIR" >&2
exit 2
fi
if [[ ! -x "$OPENPI_VENV/bin/python" ]]; then
echo "OpenPI Python venv not found: $OPENPI_VENV" >&2
exit 2
fi
if [[ ! -f "$OPENPI_REPO/scripts/serve_policy_from_checkpoint.py" ]]; then
echo "OpenPI checkout not found or missing serve script: $OPENPI_REPO" >&2
exit 2
fi
NORM_STATS_PATH="${NORM_STATS_PATH:-}"
NORM_STATS_HORIZON="${NORM_STATS_HORIZON:-}"
if [[ -z "$NORM_STATS_PATH" && -n "$NORM_STATS_HORIZON" ]]; then
case "$NORM_STATS_HORIZON" in
10|h10|H10) NORM_STATS_PATH="$CHECKPOINT_DIR/assets/norm_stats_h10.json" ;;
50|h50|H50) NORM_STATS_PATH="$CHECKPOINT_DIR/assets/norm_stats_h50.json" ;;
*)
echo "Unsupported NORM_STATS_HORIZON=$NORM_STATS_HORIZON; expected 10 or 50." >&2
exit 2
;;
esac
fi
CHECKPOINT_NORM_STATS="$CHECKPOINT_DIR/assets/norm_stats.json"
if [[ -n "$NORM_STATS_PATH" ]]; then
if [[ ! -f "$NORM_STATS_PATH" ]]; then
echo "NORM_STATS_PATH was set but does not exist: $NORM_STATS_PATH" >&2
exit 2
fi
if [[ ! -f "$CHECKPOINT_NORM_STATS" ]] || ! cmp -s "$NORM_STATS_PATH" "$CHECKPOINT_NORM_STATS"; then
mkdir -p "$(dirname "$CHECKPOINT_NORM_STATS")"
cp "$NORM_STATS_PATH" "$CHECKPOINT_NORM_STATS"
echo "Staged norm stats for serving: $CHECKPOINT_NORM_STATS"
echo " source: $NORM_STATS_PATH"
else
echo "Using checkpoint norm stats: $CHECKPOINT_NORM_STATS"
echo " source: $NORM_STATS_PATH"
fi
fi
if ! "$OPENPI_VENV/bin/python" - <<'PY'
import importlib
import sys
missing = []
for module in ("torch", "jax", "websockets", "openpi", "openpi_client", "tyro"):
try:
importlib.import_module(module)
except Exception as exc:
missing.append(f"{module}: {exc}")
if missing:
print(f"OpenPI Python venv is missing required packages: {sys.executable}", file=sys.stderr)
for item in missing:
print(f" - {item}", file=sys.stderr)
raise SystemExit(1)
PY
then
echo "Fix with: cd $OPENPI_REPO && uv sync --no-dev" >&2
exit 2
fi
TRANSFORMERS_REPLACE_DIR="$OPENPI_REPO/src/openpi/models_pytorch/transformers_replace"
if [[ -d "$TRANSFORMERS_REPLACE_DIR" ]]; then
if ! "$OPENPI_VENV/bin/python" - <<'PY' >/dev/null 2>&1
try:
from transformers.models.siglip import check
except Exception:
raise SystemExit(1)
raise SystemExit(0 if check.check_whether_transformers_replace_is_installed_correctly() else 1)
PY
then
echo "Installing OpenPI transformers replacement into: $OPENPI_VENV"
uv pip install --python "$OPENPI_VENV/bin/python" transformers==4.53.2
TRANSFORMERS_SITE="$("$OPENPI_VENV/bin/python" - <<'PY'
from pathlib import Path
import transformers
print(Path(transformers.__file__).resolve().parent)
PY
)"
cp -r "$TRANSFORMERS_REPLACE_DIR"/* "$TRANSFORMERS_SITE"/
"$OPENPI_VENV/bin/python" - <<'PY' >/dev/null
from transformers.models.siglip import check
if not check.check_whether_transformers_replace_is_installed_correctly():
raise SystemExit("OpenPI transformers replacement check still failed.")
PY
fi
fi
if ! grep -q "rtc_prev_actions" "$OPENPI_REPO/src/openpi/models_pytorch/pi0_pytorch.py" 2>/dev/null; then
echo "Warning: OpenPI RTC patch not detected in $OPENPI_REPO." >&2
echo "Run: OPENPI_REPO=$OPENPI_REPO bash deployment/openpi_rtc_patch/apply_openpi_rtc_patch.sh" >&2
fi
REFERENCE_METRICS_ARG=""
if [[ -f "$REFERENCE_METRICS_JSON" ]]; then
REFERENCE_METRICS_ARG="--reference-metrics-json $(printf '%q' "$REFERENCE_METRICS_JSON")"
else
echo "Reference inference metrics not found; FLOPs will be null until benchmarked: $REFERENCE_METRICS_JSON" >&2
fi
NORM_STATS_ARG=""
if [[ -n "$NORM_STATS_PATH" ]]; then
NORM_STATS_ARG="--norm-stats-path $(printf '%q' "$NORM_STATS_PATH")"
fi
mkdir -p "$LOG_ROOT"
OPENPI_SITE_PACKAGES="$("$OPENPI_VENV/bin/python" - <<'PY'
import site
print(site.getsitepackages()[0])
PY
)"
OPENPI_NVIDIA_LIB_PATHS="$(find "$OPENPI_SITE_PACKAGES/nvidia" -maxdepth 2 -type d -name lib 2>/dev/null | paste -sd: -)"
COMMAND=$(cat <<EOF
set -euo pipefail
cd $(printf '%q' "$OPENPI_REPO")
unset LOCAL_RANK RANK WORLD_SIZE MASTER_ADDR MASTER_PORT
export CUDA_VISIBLE_DEVICES="\${CUDA_VISIBLE_DEVICES:-0}"
export VIRTUAL_ENV=$(printf '%q' "$OPENPI_VENV")
export PATH=$(printf '%q' "$OPENPI_VENV/bin"):"\$PATH"
export LD_LIBRARY_PATH=$(printf '%q' "$OPENPI_NVIDIA_LIB_PATHS"):"\${LD_LIBRARY_PATH:-}"
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True $(printf '%q' "$OPENPI_VENV/bin/python") scripts/serve_policy_from_checkpoint.py \
--checkpoint-dir $(printf '%q' "$CHECKPOINT_DIR") \
--config-name $(printf '%q' "$CONFIG") \
--port $(printf '%q' "$PORT") \
--pytorch-device $(printf '%q' "$DEVICE") \
--num-denoise-steps $(printf '%q' "$NUM_DENOISE_STEPS") \
$NORM_STATS_ARG \
$REFERENCE_METRICS_ARG \
2>&1 | tee $(printf '%q' "$LOG_ROOT/serve_${MODEL_NAME}_${PORT}.log")
EOF
)
if [[ "$TMUX" == 1 ]]; then
env -u TMUX -u TMUX_PANE -u TMUX_TMPDIR tmux kill-session -t "$SESSION_NAME" 2>/dev/null || true
env -u TMUX -u TMUX_PANE -u TMUX_TMPDIR tmux new-session -d -s "$SESSION_NAME" bash -lc "$COMMAND"
echo "Started server tmux session: $SESSION_NAME"
echo "Attach: tmux attach -t $SESSION_NAME"
else
bash -lc "$COMMAND"
fi