""" 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()