mirror of
https://github.com/cactus-compute/needle.git
synced 2026-10-02 05:04:32 +08:00
update needle
This commit is contained in:
@@ -46,7 +46,7 @@
|
||||
## Usage
|
||||
|
||||
```
|
||||
git clone https://github.com/cactus-compute/model.git
|
||||
git clone https://github.com/cactus-compute/needle.git
|
||||
|
||||
source ./setup
|
||||
|
||||
@@ -99,15 +99,102 @@ needle [command]
|
||||
└───────────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
## Project Structure
|
||||
## TPU Factsheet
|
||||
|
||||
```
|
||||
src/
|
||||
├── model.py ······ Transformer architecture
|
||||
├── data.py ······· TinyStories loading & preprocessing
|
||||
├── train.py ······ Training loop
|
||||
├── run.py ········ Story generation from prompts
|
||||
├── test.py ······· Throughput & quality benchmarks
|
||||
├── evaluate.py ··· NLP benchmark evaluation
|
||||
└── cli.py ········ CLI entry point
|
||||
┌────────────────────┬───────────────────┬──────────────────────┬────────────────────────────┐
|
||||
│ │ v5e │ v5p │ v6e (Trillium) │
|
||||
├────────────────────┼───────────────────┼──────────────────────┼────────────────────────────┤
|
||||
│ Optimized for │ Train + inference │ Training (max perf) │ Train + inference │
|
||||
│ HBM per chip │ 16 GB │ 95 GB │ 32 GB │
|
||||
│ FLOPS (BF16) │ 197 TFLOPS │ 459 TFLOPS │ 918 TFLOPS │
|
||||
│ HBM bandwidth │ 819 GB/s │ 2,765 GB/s │ 1,640 GB/s │
|
||||
│ ICI bandwidth │ 1,600 Gbps │ 4,800 Gbps │ 3,584 Gbps │
|
||||
│ On-demand/chip/hr │ $1.20 │ $4.20 │ $2.70 │
|
||||
│ Spot/chip/hr │ $0.60 │ $2.10 │ $1.35 │
|
||||
│ Perf per $ │ 1x (baseline) │ 0.5x │ 2x │
|
||||
├────────────────────┼───────────────────┴──────────────────────┴────────────────────────────┤
|
||||
│ │ │
|
||||
│ DATASET │ │
|
||||
│ Text │ 100B tokens │
|
||||
│ Audio │ 200k × 20s = ~200M audio tokens + ~13M transcription tokens │
|
||||
│ Effective total │ ~100.5B equivalent tokens (audio has ~2-3x encoder overhead) │
|
||||
│ Storage │ ~5 GB audio (compressed) + ~400 GB text corpus │
|
||||
│ │ │
|
||||
├────────────────────┼───────────────────┬──────────────────────┬─────────────┬──────────────┤
|
||||
│ 300M multimodal │ v5litepod-4 │ v6e-4 │ v6e-8 │ v6e-16 │
|
||||
├────────────────────┼───────────────────┼──────────────────────┼─────────────┼──────────────┤
|
||||
│ Chips │ 4 │ 4 │ 8 │ 16 │
|
||||
│ Total HBM │ 64 GB │ 128 GB │ 256 GB │ 512 GB │
|
||||
│ Est. time │ ~16-21 days │ ~4-5 days │ ~2-3 days │ ~1-1.5 days │
|
||||
│ Spot $/hr │ $2.40 │ $5.40 │ $10.80 │ $21.60 │
|
||||
│ Est. total cost │ ~$900-1,200 │ ~$550-700 │ ~$550-750 │ ~$550-750 │
|
||||
└────────────────────┴───────────────────┴──────────────────────┴─────────────┴──────────────┘
|
||||
```
|
||||
|
||||
## Setup For TPU/GCP
|
||||
|
||||
- Setup gcloud 1: download the `macOS ARM` from [here](https://docs.cloud.google.com/sdk/docs/install-sdk) and uzip.
|
||||
- Setup gcloud 2: open terminal, cd to ypur downloads and run `./google-cloud-sdk/install.sh`
|
||||
- Setup gcloud 3: run `gloud init`, sign in with cactus email, should prompt for project
|
||||
- Setup gcloud 4: else, set the project with `gcloud config set project needle-488623`
|
||||
- setup gcloud 5: run `gcloud help` and read carefully
|
||||
|
||||
## TPU Guide
|
||||
|
||||
```
|
||||
needle tpu [command]
|
||||
|
||||
┌───────────────────────────────────────────────────────────────────┐
|
||||
│ │
|
||||
│ create NAME Create TPU (auto-finds zone) │
|
||||
│ --type STR Accelerator type (default: v5litepod-4) │
|
||||
│ --version STR TPU OS (default: tpu-ubuntu2204-base) │
|
||||
│ │
|
||||
│ connect NAME SSH config + first connect (auto-zone) │
|
||||
│ claude NAME Install Claude Code on instance │
|
||||
│ │
|
||||
│ stop NAME Stop instance (keeps disk) │
|
||||
│ start NAME Restart a stopped instance │
|
||||
│ delete NAME Delete instance (prompts confirmation) │
|
||||
│ list List all TPU instances │
|
||||
│ │
|
||||
│ --zone ZONE Override auto-detected zone (optional) │
|
||||
│ │
|
||||
└───────────────────────────────────────────────────────────────────┘
|
||||
|
||||
Quota increases:
|
||||
https://console.cloud.google.com/iam-admin/quotas?project=needle-488623
|
||||
```
|
||||
|
||||
## Example Workflow
|
||||
|
||||
```
|
||||
1. Create an instance (auto: finds zone → installs Claude Code → connects via SSH)
|
||||
needle tpu create my-experiment
|
||||
(exit with 'exit' or Ctrl+D)
|
||||
|
||||
2. Reconnect anytime (exit with 'exit' or Ctrl+D)
|
||||
ssh my-experiment
|
||||
or VS Code: click the '><' in the bottom left → select my-experiment
|
||||
|
||||
--- run from the instance ---
|
||||
|
||||
3. Clone the repo on your instance
|
||||
git clone https://github.com/cactus-compute/needle.git
|
||||
cd needle
|
||||
|
||||
4. Install needle
|
||||
source ./setup
|
||||
|
||||
5. Use needle as you normally would locally, like training
|
||||
needle train --wandb
|
||||
|
||||
--- back on your Mac ---
|
||||
|
||||
6. Stop when done (saves disk, stops billing)
|
||||
needle tpu stop my-experiment
|
||||
|
||||
7. (Optional) Delete instance when no longer needed
|
||||
needle tpu delete my-experiment
|
||||
```
|
||||
+45
@@ -52,6 +52,18 @@ HELP = """
|
||||
│ --benchmarks [...] wikitext2 lambada hellaswag arc_easy │
|
||||
│ --max-samples INT Samples per benchmark (default: 500) │
|
||||
│ │
|
||||
│ tpu │
|
||||
│ create NAME Create TPU (auto-finds zone) │
|
||||
│ --type STR Accelerator type (default:v5litepod-4)│
|
||||
│ --version STR TPU OS (default: tpu-ubuntu2204-base) │
|
||||
│ connect NAME SSH config + connect (auto-zone) │
|
||||
│ claude NAME Install Claude Code on instance │
|
||||
│ stop NAME Stop instance (auto-zone) │
|
||||
│ start NAME Start stopped instance (auto-zone) │
|
||||
│ delete NAME Delete instance (auto-zone) │
|
||||
│ list List all TPU instances │
|
||||
│ --zone ZONE Override auto-detected zone │
|
||||
│ │
|
||||
└───────────────────────────────────────────────────────────────────┘
|
||||
"""
|
||||
|
||||
@@ -123,6 +135,36 @@ def main():
|
||||
choices=["wikitext2", "lambada", "hellaswag", "arc_easy"])
|
||||
p.add_argument("--max-samples", type=int, default=500)
|
||||
|
||||
p = sub.add_parser("tpu", add_help=False)
|
||||
tpu_sub = p.add_subparsers(dest="tpu_action")
|
||||
|
||||
tp = tpu_sub.add_parser("create", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--type", dest="accel_type", type=str, default="v5litepod-4")
|
||||
tp.add_argument("--version", type=str, default="tpu-ubuntu2204-base")
|
||||
|
||||
tp = tpu_sub.add_parser("connect", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--zone", type=str, default=None)
|
||||
|
||||
tp = tpu_sub.add_parser("claude", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--zone", type=str, default=None)
|
||||
|
||||
tp = tpu_sub.add_parser("stop", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--zone", type=str, default=None)
|
||||
|
||||
tp = tpu_sub.add_parser("start", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--zone", type=str, default=None)
|
||||
|
||||
tp = tpu_sub.add_parser("delete", add_help=False)
|
||||
tp.add_argument("name", type=str)
|
||||
tp.add_argument("--zone", type=str, default=None)
|
||||
|
||||
tpu_sub.add_parser("list", add_help=False)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if not args.command:
|
||||
@@ -144,3 +186,6 @@ def main():
|
||||
elif args.command == "evaluate":
|
||||
from .evaluate import main as eval_main
|
||||
eval_main(args)
|
||||
elif args.command == "tpu":
|
||||
from .tpu import tpu_dispatch
|
||||
tpu_dispatch(args)
|
||||
|
||||
+228
@@ -0,0 +1,228 @@
|
||||
import getpass
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
PROJECT = "needle-488623"
|
||||
|
||||
ZONES = [
|
||||
"us-central1-a", "us-central1-b", "us-central1-f",
|
||||
"us-east1-b", "us-east1-c", "us-east1-d",
|
||||
"us-east5-a", "us-east5-b",
|
||||
"us-south1-a", "us-south1-b",
|
||||
"us-west1-a", "us-west1-b",
|
||||
"us-west4-a",
|
||||
]
|
||||
|
||||
TPU_HELP = """
|
||||
tpu commands:
|
||||
needle tpu create NAME [--type TYPE] [--version VER]
|
||||
needle tpu connect NAME [--zone ZONE]
|
||||
needle tpu claude NAME [--zone ZONE]
|
||||
needle tpu stop NAME [--zone ZONE]
|
||||
needle tpu start NAME [--zone ZONE]
|
||||
needle tpu delete NAME [--zone ZONE]
|
||||
needle tpu list
|
||||
"""
|
||||
|
||||
|
||||
def _run(cmd, check=True, capture=False, quiet=False):
|
||||
if not quiet:
|
||||
print(f"[tpu] $ {' '.join(cmd)}")
|
||||
try:
|
||||
result = subprocess.run(cmd, capture_output=capture, text=True)
|
||||
except FileNotFoundError:
|
||||
print(
|
||||
"[tpu] ERROR: 'gcloud' not found. "
|
||||
"Install: https://cloud.google.com/sdk/docs/install",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
if check and result.returncode != 0:
|
||||
if capture and result.stderr:
|
||||
print(result.stderr.strip(), file=sys.stderr)
|
||||
sys.exit(result.returncode)
|
||||
return result
|
||||
|
||||
|
||||
def _detect_zone(name):
|
||||
print(f"[tpu] Searching for '{name}'...")
|
||||
for zone in ZONES:
|
||||
result = _run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "describe", name,
|
||||
"--zone", zone, "--project", PROJECT,
|
||||
"--format", "value(name)"],
|
||||
check=False, capture=True, quiet=True,
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
print(f"[tpu] Found '{name}' in {zone}")
|
||||
return zone
|
||||
print(
|
||||
f"[tpu] ERROR: instance '{name}' not found. "
|
||||
"Run 'needle tpu list' to see instances.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def _update_ssh_config(path, host_name, new_block):
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
if os.path.exists(path):
|
||||
with open(path, "r") as f:
|
||||
content = f.read()
|
||||
pattern = rf"\n?Host {re.escape(host_name)}\n(?: .*\n)*"
|
||||
content = re.sub(pattern, "", content)
|
||||
else:
|
||||
content = ""
|
||||
with open(path, "w") as f:
|
||||
f.write(content.rstrip("\n") + "\n" if content.strip() else "")
|
||||
f.write(new_block)
|
||||
|
||||
|
||||
def tpu_create(args):
|
||||
for zone in ZONES:
|
||||
print(f"[tpu] Trying {zone}...")
|
||||
result = _run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "create", args.name,
|
||||
"--zone", zone,
|
||||
"--accelerator-type", args.accel_type,
|
||||
"--version", args.version,
|
||||
"--project", PROJECT],
|
||||
check=False, capture=True,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
print(f"[tpu] SUCCESS: created '{args.name}' in {zone}")
|
||||
args.zone = zone
|
||||
tpu_claude(args)
|
||||
tpu_connect(args)
|
||||
return
|
||||
stderr = result.stderr.strip()
|
||||
last_line = stderr.splitlines()[-1] if stderr else "unknown error"
|
||||
print(f"[tpu] {zone}: {last_line}")
|
||||
print(
|
||||
f"[tpu] ERROR: could not create '{args.name}' in any zone.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
print(
|
||||
f"[tpu] Check quota: "
|
||||
f"https://console.cloud.google.com/iam-admin/quotas?project={PROJECT}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def tpu_connect(args):
|
||||
zone = args.zone or _detect_zone(args.name)
|
||||
|
||||
result = _run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "ssh", args.name,
|
||||
"--zone", zone, "--project", PROJECT, "--dry-run"],
|
||||
capture=True,
|
||||
)
|
||||
|
||||
ip_match = re.search(
|
||||
r"(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})",
|
||||
result.stdout + result.stderr,
|
||||
)
|
||||
if not ip_match:
|
||||
print("[tpu] ERROR: could not parse IP from dry-run output.", file=sys.stderr)
|
||||
print(f"[tpu] stdout: {result.stdout}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
ip = ip_match.group(1)
|
||||
print(f"[tpu] Detected IP: {ip}")
|
||||
|
||||
ssh_config_path = os.path.expanduser("~/.ssh/config")
|
||||
user = getpass.getuser()
|
||||
block = (
|
||||
f"\nHost {args.name}\n"
|
||||
f" HostName {ip}\n"
|
||||
f" User {user}\n"
|
||||
f" IdentityFile ~/.ssh/google_compute_engine\n"
|
||||
f" CheckHostIP no\n"
|
||||
f" StrictHostKeyChecking no\n"
|
||||
)
|
||||
_update_ssh_config(ssh_config_path, args.name, block)
|
||||
print(f"[tpu] Updated {ssh_config_path} with host '{args.name}'")
|
||||
|
||||
print(f"[tpu] Connecting to {args.name} (this propagates SSH keys)...")
|
||||
subprocess.run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "ssh", args.name,
|
||||
"--zone", zone, "--project", PROJECT],
|
||||
)
|
||||
|
||||
|
||||
def tpu_stop(args):
|
||||
zone = args.zone or _detect_zone(args.name)
|
||||
_run(["gcloud", "compute", "tpus", "tpu-vm", "stop", args.name,
|
||||
"--zone", zone, "--project", PROJECT])
|
||||
print(f"[tpu] Stopped '{args.name}'")
|
||||
|
||||
|
||||
def tpu_start(args):
|
||||
zone = args.zone or _detect_zone(args.name)
|
||||
_run(["gcloud", "compute", "tpus", "tpu-vm", "start", args.name,
|
||||
"--zone", zone, "--project", PROJECT])
|
||||
print(f"[tpu] Started '{args.name}'")
|
||||
|
||||
|
||||
def tpu_delete(args):
|
||||
zone = args.zone or _detect_zone(args.name)
|
||||
answer = input(f"[tpu] Delete '{args.name}' in {zone}? This is permanent. [y/N] ")
|
||||
if answer.lower() != "y":
|
||||
print("[tpu] Aborted.")
|
||||
return
|
||||
_run(["gcloud", "compute", "tpus", "tpu-vm", "delete", args.name,
|
||||
"--zone", zone, "--project", PROJECT, "--quiet"])
|
||||
print(f"[tpu] Deleted '{args.name}'")
|
||||
|
||||
|
||||
def tpu_claude(args):
|
||||
zone = args.zone or _detect_zone(args.name)
|
||||
setup_script = (
|
||||
"curl -fsSL https://deb.nodesource.com/setup_22.x | sudo -E bash - && "
|
||||
"sudo apt-get install -y nodejs && "
|
||||
"sudo npm install -g @anthropic-ai/claude-code && "
|
||||
"echo '[tpu] Claude Code installed. Run: claude'"
|
||||
)
|
||||
print(f"[tpu] Installing Claude Code on '{args.name}'...")
|
||||
_run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "ssh", args.name,
|
||||
"--zone", zone, "--project", PROJECT,
|
||||
"--command", setup_script],
|
||||
)
|
||||
|
||||
|
||||
def tpu_list(args):
|
||||
found = False
|
||||
for zone in ZONES:
|
||||
result = _run(
|
||||
["gcloud", "compute", "tpus", "tpu-vm", "list",
|
||||
"--zone", zone, "--project", PROJECT],
|
||||
check=False, capture=True, quiet=True,
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
if not found:
|
||||
header = result.stdout.strip().splitlines()[0]
|
||||
print(f"{header} ZONE")
|
||||
found = True
|
||||
for line in result.stdout.strip().splitlines()[1:]:
|
||||
print(f"{line} {zone}")
|
||||
if not found:
|
||||
print("[tpu] No instances found.")
|
||||
|
||||
|
||||
def tpu_dispatch(args):
|
||||
actions = {
|
||||
"create": tpu_create,
|
||||
"connect": tpu_connect,
|
||||
"claude": tpu_claude,
|
||||
"stop": tpu_stop,
|
||||
"start": tpu_start,
|
||||
"delete": tpu_delete,
|
||||
"list": tpu_list,
|
||||
}
|
||||
if not args.tpu_action or args.tpu_action not in actions:
|
||||
print(TPU_HELP)
|
||||
sys.exit(0 if not args.tpu_action else 1)
|
||||
actions[args.tpu_action](args)
|
||||
Reference in New Issue
Block a user