Skip to content

Commit 2edbee5

Browse files
test(case2): add --grpc-echo-addr direct-gRPC transport for the per-block round-trip
Co-authored-by: FluffyAIcode <FluffyAIcode@users.noreply.github.com>
1 parent 2d53623 commit 2edbee5

1 file changed

Lines changed: 46 additions & 2 deletions

File tree

scripts/research/k3_specdecode_gpu_bench.py

Lines changed: 46 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -192,6 +192,7 @@ def restored_specdecode_fused(
192192
prompt, gen_tokens, block_size, device, eos_ids,
193193
block_rtt_ms: float = 0.0,
194194
sock: Optional[socket.socket] = None,
195+
grpc_call=None,
195196
) -> Dict[str, Any]:
196197
"""FUSED spec-decode engine (A+B+C): per-block O(L).
197198
@@ -269,12 +270,17 @@ def restored_specdecode_fused(
269270
# proposer's next draft block goes the other way. Round-trip the REAL
270271
# per-block payload through the socket so netem latency + serialization
271272
# + bandwidth of the actual aux tensors are all measured.
272-
if sock is not None:
273+
if sock is not None or grpc_call is not None:
273274
tn = time.perf_counter()
274275
buf = io.BytesIO()
275276
torch.save({"aux": [a.to("cpu", torch.float16) for a in new_aux],
276277
"tokens": candidate}, buf)
277-
net_bytes += _sock_roundtrip(sock, buf.getvalue())
278+
payload = buf.getvalue()
279+
if grpc_call is not None:
280+
grpc_call(payload)
281+
net_bytes += len(payload)
282+
else:
283+
net_bytes += _sock_roundtrip(sock, payload)
278284
t_network += time.perf_counter() - tn
279285
ctx_kv = drafter.extend_context_kv(
280286
ctx_kv, drafter.make_context_kv(new_aux, new_positions))
@@ -336,6 +342,10 @@ def main() -> int:
336342
help="Comma-separated netem delays (ms) to apply to "
337343
"--netem-dev between socket-mode runs (needs root + tc).")
338344
ap.add_argument("--netem-dev", default="lo")
345+
ap.add_argument("--grpc-echo-addr", default=None,
346+
help="HOST:PORT of grpc_echo_probe.py --role server. Enables "
347+
"a real gRPC per-block round-trip (direct-gRPC transport "
348+
"instead of the raw socket).")
339349
ap.add_argument("--output", default=None)
340350
args = ap.parse_args()
341351

@@ -558,6 +568,40 @@ def recall(tokens, ans):
558568
"token-level draft data plane is WAN-infeasible."),
559569
}
560570

571+
# --- Case 2 (REAL): direct-gRPC transport over the network ---
572+
if args.grpc_echo_addr:
573+
import grpc as _grpc
574+
_ident = lambda b: b # noqa: E731
575+
_opts = [("grpc.max_send_message_length", 256 * 1024 * 1024),
576+
("grpc.max_receive_message_length", 256 * 1024 * 1024),
577+
("grpc.enable_http_proxy", 0)]
578+
_ch = _grpc.insecure_channel(args.grpc_echo_addr, options=_opts)
579+
_echo = _ch.unary_unary("/echo.Echo/Echo", request_serializer=_ident,
580+
response_deserializer=_ident)
581+
_grpc.channel_ready_future(_ch).result(timeout=30)
582+
prompt0 = ids_list[0][0].tolist()
583+
r = restored_specdecode_fused(
584+
adapter, drafter, verifier, aux_layer_ids, embed_fn, lm_head_fn,
585+
prompt0, args.max_new_tokens, args.block_size, device, eos_ids,
586+
grpc_call=_echo)
587+
_ch.close()
588+
tps = r["decode_tokens_per_s"]
589+
report["crosshost_grpc_realnet"] = {
590+
"ar_baseline_tps": ar_mean,
591+
"colocated_fused_tps": fu_tps,
592+
"transport": "direct gRPC (HTTP/2) per-block round-trip",
593+
"decode_tokens_per_s": tps,
594+
"vs_ar_x": round(tps / ar_mean, 3) if ar_mean else None,
595+
"blocks": r["blocks"], "mean_accept_len": r["mean_accept_len"],
596+
"net_bytes_per_block": r["net_bytes_per_block"],
597+
"network_s": r["time_breakdown_s"]["network_rtt"],
598+
"decode_s": r["decode_s"],
599+
}
600+
print(f"[sd][grpc] direct gRPC -> {tps} tok/s "
601+
f"({report['crosshost_grpc_realnet']['vs_ar_x']}x AR, "
602+
f"{r['net_bytes_per_block']} B/block, blocks={r['blocks']})",
603+
file=sys.stderr, flush=True)
604+
561605
# --- Case 2 (REAL): two-process socket + tc netem real-network sweep ---
562606
if args.socket_echo_addr:
563607
host, port = args.socket_echo_addr.rsplit(":", 1)

0 commit comments

Comments
 (0)