jagat-primitive-org commited on
Commit
eac69bc
·
verified ·
1 Parent(s): 9362999

ship the async-MRV2 connector fix as an overlay (upstream 4e8b849)

Browse files
Files changed (1) hide show
  1. connector_mrv2.py +472 -0
connector_mrv2.py ADDED
@@ -0,0 +1,472 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3
+ """Exchange PLE data between a GPU worker and the CPU-offload process."""
4
+
5
+ import os
6
+ import queue
7
+ import threading
8
+ from dataclasses import dataclass
9
+ from multiprocessing.reduction import ForkingPickler
10
+ from typing import Any
11
+
12
+ import msgspec
13
+ import torch
14
+ import torch.nn as nn
15
+ import zmq
16
+ from cuda.bindings import driver as cuda_driver
17
+
18
+ from vllm.config import VllmConfig
19
+ from vllm.distributed.parallel_state import get_dp_group, get_tp_group
20
+ from vllm.logger import init_logger
21
+ from vllm.model_executor.layers.ple_offload_layer import (
22
+ CpuGpuSemaphore,
23
+ PleOffloadLayer,
24
+ )
25
+ from vllm.v1.ple_offload.protocol import (
26
+ PleOffloadRegistration,
27
+ PleOffloadRequest,
28
+ )
29
+
30
+ logger = init_logger(__name__)
31
+
32
+
33
+ @dataclass(frozen=True)
34
+ class _PendingPleOffloadRequest:
35
+ """Bind request metadata to its MRV2 D2H completion event."""
36
+
37
+ request: PleOffloadRequest
38
+ d2h_done_event: torch.cuda.Event | None
39
+
40
+
41
+ def _cuda_check(result: Any, operation: str) -> Any:
42
+ """Check the ``(CUresult, ...)`` tuple returned by cuda-python calls."""
43
+ error = result[0] if isinstance(result, tuple) else result
44
+ if error.value != 0:
45
+ raise RuntimeError(f"{operation} failed: {error}")
46
+ return result
47
+
48
+
49
+ class PleOffloadConnector:
50
+ """Connect a GPU runner to the shared PLE CPU worker.
51
+
52
+ MRV1 and MRV2 share the same CPU-input and CUDA-output IPC protocol.
53
+ """
54
+
55
+ def __init__(
56
+ self,
57
+ vllm_config: VllmConfig,
58
+ model: nn.Module,
59
+ device: torch.device,
60
+ ipc_addr: str,
61
+ *,
62
+ input_ids_source: torch.Tensor,
63
+ query_start_loc_source: torch.Tensor,
64
+ ngram_context_source: torch.Tensor | None,
65
+ ) -> None:
66
+ self.device = device
67
+ self.dp_rank = get_dp_group().rank_in_group
68
+ self.tp_rank = get_tp_group().rank_in_group
69
+ self._layers = self._setup_layers(vllm_config, model)
70
+
71
+ # Both runner paths stage into the same shared buffers. TP0 registers
72
+ # them with CUDA so MRV2 can use asynchronous D2H copies.
73
+ scheduler_config = vllm_config.scheduler_config
74
+ self._input_ids_buf = torch.empty(
75
+ scheduler_config.max_num_batched_tokens,
76
+ dtype=torch.int32,
77
+ device="cpu",
78
+ ).share_memory_()
79
+ self._query_start_loc_buf = torch.empty(
80
+ scheduler_config.max_num_seqs + 1,
81
+ dtype=torch.int32,
82
+ device="cpu",
83
+ ).share_memory_()
84
+ self._ngram_context_buf = None
85
+ config = vllm_config.model_config.hf_text_config
86
+ ngram_context_len = int(config.ngram_size) - 1
87
+ if ngram_context_len > 0:
88
+ self._ngram_context_buf = torch.empty(
89
+ scheduler_config.max_num_seqs,
90
+ ngram_context_len,
91
+ dtype=torch.int32,
92
+ device="cpu",
93
+ ).share_memory_()
94
+
95
+ # Runner input allocations are address-stable, so bind them once and
96
+ # pass only batch sizes through the per-forward request queue.
97
+ self._input_ids_source = input_ids_source
98
+ self._query_start_loc_source = query_start_loc_source
99
+ self._ngram_context_source = ngram_context_source
100
+ self._uses_cuda_inputs = self._input_ids_source.is_cuda
101
+ self._validate_input_sources()
102
+
103
+ self._pinned_input_buffers: list[torch.Tensor] = []
104
+ request_queue_size = (
105
+ vllm_config.max_concurrent_batches if self._uses_cuda_inputs else 1
106
+ )
107
+ self._request_queue: queue.Queue[_PendingPleOffloadRequest | None] = (
108
+ queue.Queue(maxsize=request_queue_size)
109
+ )
110
+ self._request_thread: threading.Thread | None = None
111
+ self._request_thread_ready = threading.Event()
112
+ self._zmq_ctx: zmq.Context | None = None
113
+ self._registration_socket: zmq.Socket | None = None
114
+ self._d2h_event_pool: queue.Queue[torch.cuda.Event] | None = None
115
+
116
+ try:
117
+ self._zmq_ctx = zmq.Context()
118
+ self._registration_socket = self._zmq_ctx.socket(zmq.PUSH)
119
+ self._registration_socket.connect(ipc_addr)
120
+ self._register_with_offload_worker(vllm_config, ipc_addr)
121
+
122
+ if self.tp_rank == 0:
123
+ # ForkingPickler may replace CPU storage while converting its
124
+ # sharing strategy, so register only the final addresses.
125
+ with torch.accelerator.device_index(self.device.index):
126
+ self._pin_input_buffers()
127
+ if self._uses_cuda_inputs:
128
+ self._d2h_event_pool = queue.Queue(
129
+ maxsize=vllm_config.max_concurrent_batches
130
+ )
131
+ for _ in range(vllm_config.max_concurrent_batches):
132
+ self._d2h_event_pool.put_nowait(torch.cuda.Event())
133
+ self._start_request_thread(ipc_addr)
134
+ except Exception:
135
+ self.close()
136
+ raise
137
+
138
+ def _setup_layers(
139
+ self,
140
+ vllm_config: VllmConfig,
141
+ model: nn.Module,
142
+ ) -> dict[str, PleOffloadLayer]:
143
+ """Attach output buffers and semaphores to GPU PLE placeholders."""
144
+ layers = {
145
+ name: module
146
+ for name, module in model.named_modules()
147
+ if isinstance(module, PleOffloadLayer)
148
+ }
149
+ if not layers:
150
+ raise RuntimeError(
151
+ "VLLM_PLE_CPU_OFFLOAD is enabled, but the model has no PleOffloadLayer"
152
+ )
153
+
154
+ config = vllm_config.model_config.hf_text_config
155
+ max_num_tokens = vllm_config.scheduler_config.max_num_batched_tokens
156
+ for layer in layers.values():
157
+ # The CPU worker writes results here through CUDA IPC. The GPU
158
+ # placeholder waits on the paired cross-process semaphore.
159
+ output_buffer = torch.empty(
160
+ max_num_tokens,
161
+ int(config.ple_embed_dim),
162
+ dtype=layer.get_offload_output_dtype(vllm_config.model_config.dtype),
163
+ device=self.device,
164
+ )
165
+ layer.setup_cross_process_offload(
166
+ output_buffer,
167
+ CpuGpuSemaphore(self.device),
168
+ )
169
+ return layers
170
+
171
+ def _pin_input_buffers(self) -> None:
172
+ """Page-lock shared input allocations without replacing their storage."""
173
+ buffers = [self._input_ids_buf, self._query_start_loc_buf]
174
+ if self._ngram_context_buf is not None:
175
+ buffers.append(self._ngram_context_buf)
176
+ for buffer in buffers:
177
+ if buffer.device.type != "cpu" or not buffer.is_shared():
178
+ raise RuntimeError("PLE input buffers must be shared CPU tensors")
179
+ if not buffer.is_contiguous():
180
+ raise RuntimeError("PLE input buffers must be contiguous")
181
+ _cuda_check(
182
+ cuda_driver.cuMemHostRegister(
183
+ buffer.data_ptr(),
184
+ buffer.numel() * buffer.element_size(),
185
+ cuda_driver.CU_MEMHOSTREGISTER_PORTABLE,
186
+ ),
187
+ "cuMemHostRegister(PLE input buffer)",
188
+ )
189
+ self._pinned_input_buffers.append(buffer)
190
+ if not buffer.is_pinned():
191
+ raise RuntimeError("CUDA did not page-lock a PLE input buffer")
192
+
193
+ def _unpin_input_buffers(self) -> None:
194
+ """Release CUDA registrations after the request thread has stopped."""
195
+ for buffer in reversed(self._pinned_input_buffers):
196
+ try:
197
+ _cuda_check(
198
+ cuda_driver.cuMemHostUnregister(buffer.data_ptr()),
199
+ "cuMemHostUnregister(PLE input buffer)",
200
+ )
201
+ except RuntimeError:
202
+ logger.exception("Failed to unregister a PLE input buffer")
203
+ self._pinned_input_buffers.clear()
204
+
205
+ def _register_with_offload_worker(
206
+ self, vllm_config: VllmConfig, ipc_addr: str
207
+ ) -> None:
208
+ """Register CUDA IPC outputs and shared CPU inputs with the worker."""
209
+ # Each GPU worker owns distinct output buffers, while TP0's shared
210
+ # inputs become the request source for its DP rank.
211
+ registration = PleOffloadRegistration(
212
+ worker_id=(
213
+ self.dp_rank * vllm_config.parallel_config.world_size
214
+ + vllm_config.parallel_config.rank
215
+ ),
216
+ tp_rank=self.tp_rank,
217
+ dp_rank=self.dp_rank,
218
+ gpu_output_buffers={
219
+ name: layer._gpu_output_buffer for name, layer in self._layers.items()
220
+ },
221
+ sem_flag_tensors={
222
+ name: layer._sem.flag_tensor for name, layer in self._layers.items()
223
+ },
224
+ input_ids_buf=self._input_ids_buf,
225
+ query_start_loc_buf=self._query_start_loc_buf,
226
+ ngram_context_buf=self._ngram_context_buf,
227
+ )
228
+
229
+ # ForkingPickler transmits tensors through shared-memory and CUDA IPC.
230
+ import torch.multiprocessing as torch_mp
231
+
232
+ original_strategy = torch_mp.get_sharing_strategy()
233
+ torch_mp.set_sharing_strategy("file_system")
234
+ try:
235
+ payload = ForkingPickler.dumps(registration)
236
+ finally:
237
+ torch_mp.set_sharing_strategy(original_strategy)
238
+ assert self._registration_socket is not None
239
+ self._registration_socket.send(payload)
240
+
241
+ logger.info(
242
+ "PleOffload: registered %d PleOffloadLayer(s) "
243
+ "(dp_rank=%d, tp_rank=%d, ipc_addr=%s): %s",
244
+ len(self._layers),
245
+ self.dp_rank,
246
+ self.tp_rank,
247
+ ipc_addr,
248
+ sorted(self._layers),
249
+ )
250
+
251
+ def _start_request_thread(self, ipc_addr: str) -> None:
252
+ """Start the thread that publishes batches after inputs are ready."""
253
+ self._request_thread = threading.Thread(
254
+ target=self._request_loop,
255
+ args=(ipc_addr,),
256
+ name=f"ple-offload-dp{self.dp_rank}",
257
+ daemon=True,
258
+ )
259
+ self._request_thread.start()
260
+ if not self._request_thread_ready.wait(timeout=10):
261
+ raise RuntimeError("Timed out starting the PLE request thread")
262
+
263
+ def _request_loop(self, ipc_addr: str) -> None:
264
+ """Wait for staged inputs, then notify the CPU worker."""
265
+ socket: zmq.Socket | None = None
266
+ try:
267
+ if self._zmq_ctx is None:
268
+ raise RuntimeError("PLE ZMQ context closed before thread startup")
269
+ socket = self._zmq_ctx.socket(zmq.PUSH)
270
+ socket.connect(ipc_addr)
271
+ self._request_thread_ready.set()
272
+ while True:
273
+ request = self._request_queue.get()
274
+ if request is None:
275
+ return
276
+ self._process_request(request, socket)
277
+ except Exception:
278
+ logger.exception("PLE request thread failed")
279
+ os._exit(1)
280
+ finally:
281
+ self._request_thread_ready.set()
282
+ if socket is not None:
283
+ socket.close(linger=0)
284
+
285
+ def _process_request(
286
+ self, pending: _PendingPleOffloadRequest, socket: zmq.Socket
287
+ ) -> None:
288
+ """Wait for one staged batch and publish its request."""
289
+ request = pending.request
290
+ event_pool = self._d2h_event_pool
291
+ event = pending.d2h_done_event
292
+ if self._uses_cuda_inputs:
293
+ assert event_pool is not None, "PLE D2H event pool is not initialized"
294
+ assert event is not None, "MRV2 request is missing its D2H event"
295
+ with (
296
+ torch.accelerator.device_index(self.device.index),
297
+ torch.cuda.nvtx.range("ple_offload.wait_d2h"),
298
+ ):
299
+ event.synchronize()
300
+ else:
301
+ assert pending.d2h_done_event is None
302
+ self._copy_cpu_inputs(request)
303
+
304
+ with torch.cuda.nvtx.range("ple_offload.send_request"):
305
+ socket.send(msgspec.msgpack.encode(request))
306
+ if event is not None:
307
+ assert event_pool is not None
308
+ event_pool.put_nowait(event)
309
+
310
+ def _copy_cpu_inputs(self, request: PleOffloadRequest) -> None:
311
+ """Stage MRV1's existing CPU mirrors in the notifier thread."""
312
+ num_tokens = request.num_tokens
313
+ num_reqs = request.num_reqs
314
+ with torch.cuda.nvtx.range("ple_offload.copy_input_ids"):
315
+ self._input_ids_buf[:num_tokens].copy_(self._input_ids_source[:num_tokens])
316
+ with torch.cuda.nvtx.range("ple_offload.copy_query_start_loc"):
317
+ self._query_start_loc_buf[: num_reqs + 1].copy_(
318
+ self._query_start_loc_source[: num_reqs + 1]
319
+ )
320
+ if self._ngram_context_buf is not None:
321
+ assert self._ngram_context_source is not None
322
+ with torch.cuda.nvtx.range("ple_offload.copy_ngram_context"):
323
+ self._ngram_context_buf[:num_reqs].copy_(
324
+ self._ngram_context_source[:num_reqs]
325
+ )
326
+
327
+ def _validate_input_sources(self) -> None:
328
+ """Validate fixed runner sources against shared input buffers."""
329
+ sources = [
330
+ ("input_ids", self._input_ids_source, self._input_ids_buf),
331
+ (
332
+ "query_start_loc",
333
+ self._query_start_loc_source,
334
+ self._query_start_loc_buf,
335
+ ),
336
+ ]
337
+ if (self._ngram_context_source is None) != (self._ngram_context_buf is None):
338
+ raise ValueError("PLE ngram_context source and buffer must match")
339
+ if self._ngram_context_source is not None:
340
+ assert self._ngram_context_buf is not None
341
+ sources.append(
342
+ (
343
+ "ngram_context",
344
+ self._ngram_context_source,
345
+ self._ngram_context_buf,
346
+ )
347
+ )
348
+
349
+ expected_device = self.device if self._uses_cuda_inputs else torch.device("cpu")
350
+ for name, source, buffer in sources:
351
+ if (
352
+ source.device != expected_device
353
+ or source.dtype != buffer.dtype
354
+ or source.ndim != buffer.ndim
355
+ or source.shape[0] < buffer.shape[0]
356
+ or source.shape[1:] != buffer.shape[1:]
357
+ ):
358
+ raise ValueError(f"PLE {name} source is incompatible")
359
+
360
+ def _enqueue_cuda_inputs(
361
+ self,
362
+ request: PleOffloadRequest,
363
+ d2h_done_event: torch.cuda.Event,
364
+ ) -> None:
365
+ """Stage MRV2 inputs on the model stream and record completion."""
366
+ with torch.accelerator.device_index(self.device.index):
367
+ stream = torch.cuda.current_stream(self.device)
368
+ with torch.cuda.nvtx.range("ple_offload.copy_input_ids"):
369
+ self._input_ids_buf[: request.num_tokens].copy_(
370
+ self._input_ids_source[: request.num_tokens],
371
+ non_blocking=True,
372
+ )
373
+ with torch.cuda.nvtx.range("ple_offload.copy_query_start_loc"):
374
+ self._query_start_loc_buf[: request.num_reqs + 1].copy_(
375
+ self._query_start_loc_source[: request.num_reqs + 1],
376
+ non_blocking=True,
377
+ )
378
+ if self._ngram_context_buf is not None:
379
+ assert self._ngram_context_source is not None
380
+ with torch.cuda.nvtx.range("ple_offload.copy_ngram_context"):
381
+ self._ngram_context_buf[: request.num_reqs].copy_(
382
+ self._ngram_context_source[: request.num_reqs],
383
+ non_blocking=True,
384
+ )
385
+ d2h_done_event.record(stream)
386
+
387
+ def _launch(
388
+ self,
389
+ num_reqs: int,
390
+ num_tokens: int,
391
+ ) -> None:
392
+ """Stage or queue one batch for request publication."""
393
+ # Inputs are replicated across TP ranks. One request per DP rank drives
394
+ # the CPU result fan-out to every registered TP output buffer.
395
+ if self.tp_rank != 0:
396
+ return
397
+
398
+ request = PleOffloadRequest(
399
+ dp_rank=self.dp_rank,
400
+ num_tokens=num_tokens,
401
+ num_reqs=num_reqs,
402
+ )
403
+ d2h_done_event = None
404
+ if self._uses_cuda_inputs:
405
+ assert self._d2h_event_pool is not None, (
406
+ "PLE D2H event pool is not initialized"
407
+ )
408
+ try:
409
+ d2h_done_event = self._d2h_event_pool.get_nowait()
410
+ except queue.Empty as exc:
411
+ raise RuntimeError(
412
+ "PLE has more MRV2 requests than configured concurrent batches"
413
+ ) from exc
414
+ self._enqueue_cuda_inputs(request, d2h_done_event)
415
+ self._request_queue.put_nowait(
416
+ _PendingPleOffloadRequest(request, d2h_done_event)
417
+ )
418
+
419
+ def prepare_forward(
420
+ self,
421
+ num_reqs: int,
422
+ num_tokens: int,
423
+ dummy_run: bool,
424
+ ) -> None:
425
+ """Submit real inputs or satisfy the PLE wait for a dummy forward."""
426
+ if dummy_run:
427
+ self.signal_dummy_outputs(num_tokens)
428
+ return
429
+ self._launch(num_reqs, num_tokens)
430
+
431
+ def signal_dummy_outputs(self, num_tokens: int) -> None:
432
+ """Locally satisfy PLE waits for dummy and capture forwards."""
433
+ # Dummy and capture forwards do not send CPU requests, but every PLE
434
+ # placeholder still waits for a completed output semaphore.
435
+ stream = torch.cuda.current_stream(self.device)
436
+ for layer in self._layers.values():
437
+ layer._gpu_output_buffer[:num_tokens].zero_()
438
+ layer._sem.signal(stream)
439
+
440
+ def release_outputs(self) -> None:
441
+ """Mark GPU output buffers reusable after the model consumes them."""
442
+ # Reset only after the consumer forward so the CPU worker cannot
443
+ # overwrite an output that a GPU PLE placeholder may still read.
444
+ stream = torch.cuda.current_stream(self.device)
445
+ for layer in self._layers.values():
446
+ layer.release_offloaded_output(stream)
447
+
448
+ def close(self) -> None:
449
+ """Stop request transport and release host registrations."""
450
+ request_thread = self._request_thread
451
+ if request_thread is not None and request_thread.is_alive():
452
+ try:
453
+ self._request_queue.put(None, timeout=5)
454
+ except queue.Full:
455
+ logger.error("Timed out stopping the PLE request thread")
456
+ request_thread.join(timeout=5)
457
+ if request_thread is not None and request_thread.is_alive():
458
+ # The thread may still access the registered buffers or ZMQ context.
459
+ logger.error("PLE request thread did not stop; deferring resource cleanup")
460
+ return
461
+ self._request_thread = None
462
+
463
+ if self._pinned_input_buffers:
464
+ with torch.accelerator.device_index(self.device.index):
465
+ self._unpin_input_buffers()
466
+ self._d2h_event_pool = None
467
+ if self._registration_socket is not None:
468
+ self._registration_socket.close(linger=0)
469
+ self._registration_socket = None
470
+ if self._zmq_ctx is not None:
471
+ self._zmq_ctx.term()
472
+ self._zmq_ctx = None