@@ -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
0 commit comments