mirror of
https://github.com/cactus-compute/needle.git
synced 2026-10-02 05:04:32 +08:00
Refactor project structure and update dependencies
- Removed GCS utilities and TPU management scripts from the codebase. - Updated project name to "cactus-needle" and version to "2.0.0" in pyproject.toml. - Simplified setup scripts for virtual environment creation and dependency installation. - Removed unnecessary dependencies from requirements.txt and setup scripts. - Enhanced compatibility with Python 3.9 and above. - Updated README and documentation references accordingly.
This commit is contained in:
@@ -1,129 +1,30 @@
|
||||
#!/usr/bin/env bash
|
||||
set -e
|
||||
|
||||
VENV_DIR=".venv"
|
||||
|
||||
# On Ubuntu/Debian, ensure Python >= 3.11 and venv are available
|
||||
if [ -f /etc/os-release ] && grep -qi 'ubuntu\|debian' /etc/os-release; then
|
||||
PYTHON_OK=$(python3 -c "import sys; print(int(sys.version_info >= (3, 11)))" 2>/dev/null || echo "0")
|
||||
if [ "$PYTHON_OK" = "0" ]; then
|
||||
echo "Python >= 3.11 required. Installing..."
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq software-properties-common
|
||||
sudo add-apt-repository -y ppa:deadsnakes/ppa
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq python3.11 python3.11-venv python3.11-dev
|
||||
PYTHON=python3.11
|
||||
else
|
||||
PYTHON=python3
|
||||
# Ensure venv package is installed
|
||||
if ! $PYTHON -m venv --help &>/dev/null; then
|
||||
echo "Installing python3-venv..."
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq python3-venv
|
||||
fi
|
||||
fi
|
||||
else
|
||||
# On macOS or other systems, find Python >= 3.11
|
||||
PYTHON=""
|
||||
for candidate in python3.14 python3.13 python3.12 python3.11 python3; do
|
||||
if command -v "$candidate" &>/dev/null; then
|
||||
PY_OK=$("$candidate" -c "import sys; print(int(sys.version_info >= (3, 11)))" 2>/dev/null || echo "0")
|
||||
if [ "$PY_OK" = "1" ]; then
|
||||
PYTHON="$candidate"
|
||||
break
|
||||
fi
|
||||
fi
|
||||
done
|
||||
if [ -z "$PYTHON" ]; then
|
||||
echo "Error: Python >= 3.11 is required. Install it with: brew install python@3.12"
|
||||
return 1 2>/dev/null || exit 1
|
||||
PYTHON=""
|
||||
for candidate in python3.12 python3.11 python3.10 python3.9 python3; do
|
||||
if command -v "$candidate" >/dev/null 2>&1 && \
|
||||
"$candidate" -c "import sys; exit(0 if sys.version_info >= (3, 9) else 1)" 2>/dev/null; then
|
||||
PYTHON="$candidate"
|
||||
break
|
||||
fi
|
||||
done
|
||||
if [ -z "$PYTHON" ]; then
|
||||
echo "Python >= 3.9 is required." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Recreate venv if it doesn't exist or was built with the wrong Python
|
||||
RECREATE=0
|
||||
if [ ! -d "$VENV_DIR" ]; then
|
||||
RECREATE=1
|
||||
elif [ ! -f "$VENV_DIR/bin/activate" ]; then
|
||||
RECREATE=1
|
||||
elif [ -f "$VENV_DIR/bin/python3" ]; then
|
||||
VENV_PY_OK=$("$VENV_DIR/bin/python3" -c "import sys; print(int(sys.version_info >= (3, 11)))" 2>/dev/null || echo "0")
|
||||
if [ "$VENV_PY_OK" = "0" ]; then
|
||||
RECREATE=1
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$RECREATE" = "1" ]; then
|
||||
echo "Creating virtual environment with $PYTHON..."
|
||||
rm -rf "$VENV_DIR"
|
||||
$PYTHON -m venv "$VENV_DIR"
|
||||
else
|
||||
echo "Virtual environment already exists."
|
||||
fi
|
||||
|
||||
[ -d "$VENV_DIR" ] || "$PYTHON" -m venv "$VENV_DIR"
|
||||
source "$VENV_DIR/bin/activate"
|
||||
|
||||
echo "Installing dependencies..."
|
||||
pip install --upgrade pip -q
|
||||
pip install -e . -q
|
||||
|
||||
# Detect accelerator and install the matching JAX build
|
||||
ACCELERATOR="cpu"
|
||||
if command -v nvidia-smi &>/dev/null && nvidia-smi &>/dev/null; then
|
||||
ACCELERATOR="gpu"
|
||||
elif ls /dev/accel* &>/dev/null 2>&1 || [ -n "$TPU_NAME" ] || [ -d /sys/class/accel ]; then
|
||||
ACCELERATOR="tpu"
|
||||
fi
|
||||
|
||||
case "$ACCELERATOR" in
|
||||
gpu)
|
||||
echo "Detected NVIDIA GPU. Installing jax[cuda12]..."
|
||||
pip install -U "jax[cuda12]"
|
||||
# GPU-friendly defaults: silence XLA autotuner spam, cache compiled
|
||||
# kernels across runs, and let JAX claim most of HBM.
|
||||
export TF_CPP_MIN_LOG_LEVEL=2
|
||||
export JAX_COMPILATION_CACHE_DIR="$HOME/.cache/jax"
|
||||
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.95
|
||||
mkdir -p "$JAX_COMPILATION_CACHE_DIR"
|
||||
;;
|
||||
tpu)
|
||||
echo "Detected TPU. Installing jax[tpu]..."
|
||||
pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
|
||||
# Load TPU kernel modules if not already loaded
|
||||
if ! ls /dev/accel* &>/dev/null 2>&1; then
|
||||
echo "Loading TPU kernel modules..."
|
||||
for mod in gasket tpu_common tpu_v2 tpu_v3 tpu_v4 tpu_v4_lite; do
|
||||
sudo modprobe "$mod" 2>/dev/null || true
|
||||
done
|
||||
fi
|
||||
if ls /dev/accel* &>/dev/null 2>&1; then
|
||||
echo "TPU devices: $(ls /dev/accel* 2>/dev/null | wc -w) chip(s)"
|
||||
fi
|
||||
;;
|
||||
*)
|
||||
echo "No GPU or TPU detected. Installing CPU-only jax..."
|
||||
pip install -U jax
|
||||
;;
|
||||
esac
|
||||
|
||||
# Sanity check: report the JAX backend that was actually selected
|
||||
python -c "import jax; print('JAX devices:', jax.devices())" 2>/dev/null || \
|
||||
echo "Warning: jax import failed; check the install above."
|
||||
|
||||
if [ -f /sys/kernel/mm/transparent_hugepage/enabled ]; then
|
||||
echo "Enabling transparent hugepages..."
|
||||
sudo sh -c "echo always > /sys/kernel/mm/transparent_hugepage/enabled" 2>/dev/null || true
|
||||
fi
|
||||
|
||||
if [ -z "$WANDB_API_KEY" ]; then
|
||||
echo ""
|
||||
echo " Enter your W&B API key (https://wandb.ai/authorize)"
|
||||
echo " Press Enter to skip."
|
||||
printf " > "
|
||||
read WANDB_API_KEY || true
|
||||
if [ -n "$WANDB_API_KEY" ]; then
|
||||
export WANDB_API_KEY
|
||||
fi
|
||||
if command -v nvidia-smi >/dev/null 2>&1 && nvidia-smi >/dev/null 2>&1; then
|
||||
echo "NVIDIA GPU detected; installing jax[cuda12]..."
|
||||
pip install -U "jax[cuda12]" -q
|
||||
fi
|
||||
|
||||
needle --help
|
||||
|
||||
Reference in New Issue
Block a user