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