Spaces:
Paused
Paused
File size: 4,493 Bytes
6abf0e6 cef3c52 6abf0e6 5ec0cf1 6abf0e6 5ec0cf1 6abf0e6 5ec0cf1 6abf0e6 1ec7a72 6abf0e6 | 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 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | """
Colab Node Runner script for ShardFlow.
Usage in Google Colab:
1. !pip install -q torch transformers tokenizers safetensors accelerate fastapi uvicorn requests pydantic sse-starlette
2. !git clone https://github.com/rautaditya2606/Shardflow.git /content/Shardflow && cd /content/Shardflow && pip install -e .
3. !python scripts/colab_runner.py --registry-url https://your-registry-url.onrender.com --model meta-llama/Meta-Llama-3-8B
"""
import argparse
import asyncio
import logging
import time
import requests
import torch
from shardflow.transport.tunnel import start_cloudflare_tcp_tunnel, start_bore_tunnel
from shardflow.node.layer_loader import load_layer_slice
from shardflow.node.node import PipelineNode
logger = logging.getLogger("shardflow.colab_runner")
def main():
parser = argparse.ArgumentParser(description="ShardFlow Colab Node Runner")
parser.add_argument("--registry-url", required=True, help="Registry URL (e.g. https://shardflow-v0-1-0.onrender.com)")
parser.add_argument("--model", default="Qwen/Qwen2.5-14B-Instruct", help="Model path or HF model ID")
parser.add_argument("--port", type=int, default=9500, help="Local TCP port")
parser.add_argument("--tunnel", choices=["bore", "cloudflare"], default="bore", help="Tunnel backend (default: bore)")
parser.add_argument("--node-id", default=None, help="Unique node identifier")
args = parser.parse_args()
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(name)s] %(levelname)s: %(message)s",
)
node_id = args.node_id or f"colab-node-{int(time.time())}"
local_port = args.port
if args.tunnel == "bore":
logger.info("Starting bore tunnel on local port %d...", local_port)
tunnel_proc, pub_host, pub_port = start_bore_tunnel(local_port)
else:
logger.info("Starting Cloudflare TCP tunnel on local port %d...", local_port)
tunnel_proc, pub_host, pub_port = start_cloudflare_tcp_tunnel(local_port)
logger.info("Tunnel established at %s:%d", pub_host, pub_port)
vram = 0.0
if torch.cuda.is_available():
vram = torch.cuda.get_device_properties(0).total_memory / (1024 * 1024)
logger.info("Registering node %s with registry %s...", node_id, args.registry_url)
reg_payload = {
"node_id": node_id,
"addr": pub_host,
"port": pub_port,
"vram_available_mb": vram,
"vram_total_mb": vram,
"model_id": args.model,
}
resp = requests.post(f"{args.registry_url.rstrip('/')}/register", json=reg_payload, timeout=15.0)
resp.raise_for_status()
assignment = resp.json()
layer_start = assignment["layer_start"]
layer_end = assignment["layer_end"]
is_first = assignment["is_first_node"]
is_last = assignment["is_last_node"]
next_host = assignment.get("next_node_host")
next_port = assignment.get("next_node_port")
logger.info(
"Successfully registered! Assigned layers [%d, %d) (is_first=%s, is_last=%s)",
layer_start, layer_end, is_first, is_last
)
if next_host:
logger.info("Next node routing target: %s:%d", next_host, next_port)
device = "cuda" if torch.cuda.is_available() else "cpu"
logger.info("Loading layer slice [%d, %d) onto device %s...", layer_start, layer_end, device)
model_slice = load_layer_slice(
model_path=args.model,
layer_start=layer_start,
layer_end=layer_end,
include_norm=is_last,
include_lm_head=is_last,
device=device,
)
node = PipelineNode(
model_slice=model_slice,
is_first_node=is_first,
is_last_node=is_last,
next_node_host=next_host,
next_node_port=next_port,
listen_host="0.0.0.0",
listen_port=local_port,
)
async def heartbeat_loop():
"""Periodically ping Topology Registry to keep node alive."""
hb_url = f"{args.registry_url.rstrip('/')}/heartbeat"
hb_payload = {"node_id": node_id}
while True:
await asyncio.sleep(10.0)
try:
requests.post(hb_url, json=hb_payload, timeout=5.0)
except Exception as e:
logger.debug("Heartbeat ping error: %s", e)
async def run_node():
asyncio.create_task(heartbeat_loop())
await node.serve_forever()
logger.info("Pipeline node running with background heartbeat...")
asyncio.run(run_node())
if __name__ == "__main__":
main()
|