Upload UR5 full fine-tuned checkpoint: pi05_pour_full configs/deployment_serve_policy.sh
4bf51a3 verified | 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 | |