Skip to content

Commit 76efd95

Browse files
committed
Address VMM transaction review feedback
1 parent 98265d6 commit 76efd95

4 files changed

Lines changed: 75 additions & 564 deletions

File tree

‎cuda_core/cuda/core/_memory/_virtual_memory_resource.py‎

Lines changed: 17 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -317,28 +317,19 @@ def _grow_allocation_fast_path(
317317
"""
318318
with Transaction() as trans:
319319
# Create new physical memory for the additional size
320-
trans.append(
320+
trans.on_failure(
321321
lambda np=new_ptr, s=aligned_additional_size: raise_if_driver_error(driver.cuMemAddressFree(np, s)[0])
322322
)
323323
res, new_handle = driver.cuMemCreate(aligned_additional_size, prop, 0)
324324
raise_if_driver_error(res)
325-
new_handle_released = False
326-
327-
def _release_new_handle() -> None:
328-
nonlocal new_handle_released
329-
if not new_handle_released:
330-
raise_if_driver_error(driver.cuMemRelease(new_handle)[0])
331-
new_handle_released = True
332-
333-
# Register undo for creation. Callback is conditional to avoid
334-
# double-release after an explicit successful release.
335-
trans.append(_release_new_handle)
325+
trans.on_failure(lambda h=new_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
326+
trans.on_success(lambda h=new_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
336327

337328
# Map the new physical memory to the extended VA range
338329
(res,) = driver.cuMemMap(new_ptr, aligned_additional_size, 0, new_handle, 0)
339330
raise_if_driver_error(res)
340331
# Register undo for mapping
341-
trans.append(
332+
trans.on_failure(
342333
lambda np=new_ptr, s=aligned_additional_size: raise_if_driver_error(driver.cuMemUnmap(np, s)[0])
343334
)
344335

@@ -348,9 +339,6 @@ def _release_new_handle() -> None:
348339
(res,) = driver.cuMemSetAccess(new_ptr, aligned_additional_size, descs, len(descs))
349340
raise_if_driver_error(res)
350341

351-
# Release handle ownership now that mapping is stable.
352-
_release_new_handle()
353-
354342
# All succeeded, cancel undo actions
355343
trans.commit()
356344

@@ -394,24 +382,15 @@ def _grow_allocation_slow_path(
394382
res, new_ptr = driver.cuMemAddressReserve(total_aligned_size, addr_align, 0, 0)
395383
raise_if_driver_error(res)
396384
# Register undo for VA reservation
397-
trans.append(
385+
trans.on_failure(
398386
lambda np=new_ptr, s=total_aligned_size: raise_if_driver_error(driver.cuMemAddressFree(np, s)[0])
399387
)
400388

401389
# Get the old allocation handle for remapping
402390
result, old_handle = driver.cuMemRetainAllocationHandle(buf.handle)
403391
raise_if_driver_error(result)
404-
old_handle_released = False
405-
406-
def _release_old_handle() -> None:
407-
nonlocal old_handle_released
408-
if not old_handle_released:
409-
raise_if_driver_error(driver.cuMemRelease(old_handle)[0])
410-
old_handle_released = True
411-
412-
# Register undo for old handle. Callback is conditional to avoid
413-
# double-release after explicit success.
414-
trans.append(_release_old_handle)
392+
trans.on_failure(lambda h=old_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
393+
trans.on_success(lambda h=old_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
415394

416395
# Unmap the old VA range (aligned previous size)
417396
aligned_prev_size = total_aligned_size - aligned_additional_size
@@ -427,37 +406,27 @@ def _remap_old() -> None:
427406
# TODO: consider logging this exception
428407
pass
429408

430-
trans.append(_remap_old)
409+
trans.on_failure(_remap_old)
431410

432411
# Remap the old physical memory to the new VA range (aligned previous size)
433412
(res,) = driver.cuMemMap(int(new_ptr), aligned_prev_size, 0, old_handle, 0)
434413
raise_if_driver_error(res)
435414

436415
# Register undo for mapping
437-
trans.append(lambda np=new_ptr, s=aligned_prev_size: raise_if_driver_error(driver.cuMemUnmap(np, s)[0]))
416+
trans.on_failure(lambda np=new_ptr, s=aligned_prev_size: raise_if_driver_error(driver.cuMemUnmap(np, s)[0]))
438417

439418
# Create new physical memory for the additional size
440419
res, new_handle = driver.cuMemCreate(aligned_additional_size, prop, 0)
441420
raise_if_driver_error(res)
442-
443-
new_handle_released = False
444-
445-
def _release_new_handle() -> None:
446-
nonlocal new_handle_released
447-
if not new_handle_released:
448-
raise_if_driver_error(driver.cuMemRelease(new_handle)[0])
449-
new_handle_released = True
450-
451-
# Register undo for new physical memory. Callback is conditional to
452-
# avoid double-release after explicit success.
453-
trans.append(_release_new_handle)
421+
trans.on_failure(lambda h=new_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
422+
trans.on_success(lambda h=new_handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
454423

455424
# Map the new physical memory to the extended portion (aligned offset)
456425
(res,) = driver.cuMemMap(int(new_ptr) + aligned_prev_size, aligned_additional_size, 0, new_handle, 0)
457426
raise_if_driver_error(res)
458427

459428
# Register undo for mapping
460-
trans.append(
429+
trans.on_failure(
461430
lambda base=int(new_ptr), offs=aligned_prev_size, s=aligned_additional_size: raise_if_driver_error(
462431
driver.cuMemUnmap(base + offs, s)[0]
463432
)
@@ -469,10 +438,6 @@ def _release_new_handle() -> None:
469438
(res,) = driver.cuMemSetAccess(new_ptr, total_aligned_size, descs, len(descs))
470439
raise_if_driver_error(res)
471440

472-
# Release handles once all operations that need them have completed.
473-
_release_new_handle()
474-
_release_old_handle()
475-
476441
# All succeeded, cancel undo actions
477442
trans.commit()
478443

@@ -576,28 +541,20 @@ def allocate(self, size: int, *, stream: Stream | GraphBuilder | None = None) ->
576541
# ---- Create physical memory ----
577542
res, handle = driver.cuMemCreate(aligned_size, prop, 0)
578543
raise_if_driver_error(res)
579-
handle_released = False
580-
581-
def _release_handle() -> None:
582-
nonlocal handle_released
583-
if not handle_released:
584-
raise_if_driver_error(driver.cuMemRelease(handle)[0])
585-
handle_released = True
586-
587-
# Register undo for physical memory. Callback is conditional to
588-
# avoid double-release after explicit success.
589-
trans.append(_release_handle)
544+
trans.on_failure(lambda h=handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
545+
# Once mapped, the physical allocation is kept alive without the creation reference.
546+
trans.on_success(lambda h=handle: raise_if_driver_error(driver.cuMemRelease(h)[0]))
590547

591548
# ---- Reserve VA space ----
592549
# Potentially, use a separate size for the VA reservation from the physical allocation size
593550
res, ptr = driver.cuMemAddressReserve(aligned_size, addr_align, config.addr_hint, 0)
594551
raise_if_driver_error(res)
595552
# Register undo for VA reservation
596-
trans.append(lambda p=ptr, s=aligned_size: raise_if_driver_error(driver.cuMemAddressFree(p, s)[0]))
553+
trans.on_failure(lambda p=ptr, s=aligned_size: raise_if_driver_error(driver.cuMemAddressFree(p, s)[0]))
597554

598555
# ---- Map physical memory into VA ----
599556
(res,) = driver.cuMemMap(ptr, aligned_size, 0, handle, 0)
600-
trans.append(lambda p=ptr, s=aligned_size: raise_if_driver_error(driver.cuMemUnmap(p, s)[0]))
557+
trans.on_failure(lambda p=ptr, s=aligned_size: raise_if_driver_error(driver.cuMemUnmap(p, s)[0]))
601558
raise_if_driver_error(res)
602559

603560
# ---- Set access for owner + peers ----
@@ -606,9 +563,6 @@ def _release_handle() -> None:
606563
(res,) = driver.cuMemSetAccess(ptr, aligned_size, descs, len(descs))
607564
raise_if_driver_error(res)
608565

609-
# Release handle ownership once map+access setup succeeded.
610-
_release_handle()
611-
612566
trans.commit()
613567

614568
# Done — return a Buffer that tracks this VA range

‎cuda_core/cuda/core/_utils/cuda_utils.pyi‎

Lines changed: 18 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -19,21 +19,22 @@ class NVRTCError(CUDAError):
1919

2020
class Transaction:
2121
"""
22-
A context manager for transactional operations with undo capability.
22+
A context manager for transactional operations with failure and success callbacks.
2323
24-
The Transaction class allows you to register undo actions (callbacks) that will be executed
25-
if the transaction is not committed before exiting the context. This is useful for managing
26-
resources or operations that need to be rolled back in case of errors or early exits.
24+
Failure callbacks are executed in LIFO order if the transaction exits without being committed.
25+
Success callbacks are executed in FIFO order when the transaction is committed.
2726
2827
Usage:
2928
with Transaction() as txn:
30-
txn.append(some_cleanup_function, arg1, arg2)
29+
txn.on_failure(some_cleanup_function, arg1, arg2)
30+
txn.on_success(some_finalize_function, arg1, arg2)
3131
# ... perform operations ...
32-
txn.commit() # Disarm undo actions; nothing will be rolled back on exit
32+
txn.commit()
3333
3434
Methods:
35-
append(fn, *args, **kwargs): Register an undo action to be called on rollback.
36-
commit(): Disarm all undo actions; nothing will be rolled back on exit.
35+
on_failure(fn, *args, **kwargs): Register a callback to be called on rollback.
36+
on_success(fn, *args, **kwargs): Register a callback to be called on commit.
37+
commit(): Disarm failure callbacks and run success callbacks.
3738
"""
3839

3940
def __init__(self) -> None:
@@ -45,15 +46,21 @@ class Transaction:
4546
def __exit__(self, exc_type, exc, tb):
4647
...
4748

48-
def append(self, fn: Callable[..., Any], /, *args: Any, **kwargs) -> None:
49+
def on_failure(self, fn: Callable[..., Any], /, *args: Any, **kwargs) -> None:
4950
"""
50-
Register an undo action (runs if the with-block exits without commit()).
51+
Register a failure callback (runs if the with-block exits without commit()).
52+
Values are bound now via partial so late mutations don't bite you.
53+
"""
54+
55+
def on_success(self, fn: Callable[..., Any], /, *args: Any, **kwargs) -> None:
56+
"""
57+
Register a success callback (runs in FIFO order during commit()).
5158
Values are bound now via partial so late mutations don't bite you.
5259
"""
5360

5461
def commit(self) -> None:
5562
"""
56-
Disarm all undo actions. After this, exiting the with-block does nothing.
63+
Disarm all failure callbacks, then run success callbacks in FIFO order.
5764
"""
5865
_keep_driver_in_stub: 'driver.CUresult'
5966
_keep_nvrtc_in_stub: 'nvrtc.nvrtcResult'

‎cuda_core/cuda/core/_utils/cuda_utils.pyx‎

Lines changed: 27 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -306,24 +306,26 @@ def is_nested_sequence(obj: object) -> bool:
306306

307307
class Transaction:
308308
"""
309-
A context manager for transactional operations with undo capability.
309+
A context manager for transactional operations with failure and success callbacks.
310310

311-
The Transaction class allows you to register undo actions (callbacks) that will be executed
312-
if the transaction is not committed before exiting the context. This is useful for managing
313-
resources or operations that need to be rolled back in case of errors or early exits.
311+
Failure callbacks are executed in LIFO order if the transaction exits without being committed.
312+
Success callbacks are executed in FIFO order when the transaction is committed.
314313

315314
Usage:
316315
with Transaction() as txn:
317-
txn.append(some_cleanup_function, arg1, arg2)
316+
txn.on_failure(some_cleanup_function, arg1, arg2)
317+
txn.on_success(some_finalize_function, arg1, arg2)
318318
# ... perform operations ...
319-
txn.commit() # Disarm undo actions; nothing will be rolled back on exit
319+
txn.commit()
320320

321321
Methods:
322-
append(fn, *args, **kwargs): Register an undo action to be called on rollback.
323-
commit(): Disarm all undo actions; nothing will be rolled back on exit.
322+
on_failure(fn, *args, **kwargs): Register a callback to be called on rollback.
323+
on_success(fn, *args, **kwargs): Register a callback to be called on commit.
324+
commit(): Disarm failure callbacks and run success callbacks.
324325
"""
325326
def __init__(self) -> None:
326327
self._stack = ExitStack()
328+
self._on_success: list[Callable[[], Any]] = []
327329
self._entered = False
328330
329331
def __enter__(self):
@@ -334,23 +336,36 @@ class Transaction:
334336
def __exit__(self, exc_type, exc, tb):
335337
# If exit callbacks remain, they'll run in LIFO order.
336338
self._entered = False
339+
self._on_success.clear()
337340
return self._stack.__exit__(exc_type, exc, tb)
338341
339-
def append(self, fn: Callable[..., Any], /, *args: Any, **kwargs: Any) -> None:
342+
def on_failure(self, fn: Callable[..., Any], /, *args: Any, **kwargs: Any) -> None:
340343
"""
341-
Register an undo action (runs if the with-block exits without commit()).
344+
Register a failure callback (runs if the with-block exits without commit()).
342345
Values are bound now via partial so late mutations don't bite you.
343346
"""
344347
if not self._entered:
345-
raise RuntimeError("Transaction must be entered before append()")
348+
raise RuntimeError("Transaction must be entered before on_failure()")
346349
self._stack.callback(partial(fn, *args, **kwargs))
347350
351+
def on_success(self, fn: Callable[..., Any], /, *args: Any, **kwargs: Any) -> None:
352+
"""
353+
Register a success callback (runs in FIFO order during commit()).
354+
Values are bound now via partial so late mutations don't bite you.
355+
"""
356+
if not self._entered:
357+
raise RuntimeError("Transaction must be entered before on_success()")
358+
self._on_success.append(partial(fn, *args, **kwargs))
359+
348360
def commit(self) -> None:
349361
"""
350-
Disarm all undo actions. After this, exiting the with-block does nothing.
362+
Disarm all failure callbacks, then run success callbacks in FIFO order.
351363
"""
352364
# pop_all() empties this stack so no callbacks are triggered on exit.
353365
self._stack.pop_all()
366+
for fn in self._on_success:
367+
fn()
368+
self._on_success.clear()
354369
355370
356371
# Track whether we've already warned about fork method

0 commit comments

Comments
 (0)