Shardflow / scripts /colab_runner.py
rautaditya2606's picture
Configure clean Gradio HF Space deployment
cef3c52
Raw
History Blame Contribute Delete
4.49 kB
"""
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()