@@ -329,19 +329,18 @@ def _grow_allocation_fast_path(
329329 """
330330 with Transaction () as trans :
331331 # Create new physical memory for the additional size
332- trans .append (
332+ trans .on_failure (
333333 lambda np = new_ptr , s = aligned_additional_size : raise_if_driver_error (driver .cuMemAddressFree (np , s )[0 ])
334334 )
335335 res , new_handle = driver .cuMemCreate (aligned_additional_size , prop , 0 )
336336 raise_if_driver_error (res )
337- # Register undo for creation
338- trans .append (lambda h = new_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
337+ trans .on_exit (lambda h = new_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
339338
340339 # Map the new physical memory to the extended VA range
341340 (res ,) = driver .cuMemMap (new_ptr , aligned_additional_size , 0 , new_handle , 0 )
342341 raise_if_driver_error (res )
343342 # Register undo for mapping
344- trans .append (
343+ trans .on_failure (
345344 lambda np = new_ptr , s = aligned_additional_size : raise_if_driver_error (driver .cuMemUnmap (np , s )[0 ])
346345 )
347346
@@ -393,15 +392,14 @@ def _grow_allocation_slow_path(
393392 res , new_ptr = driver .cuMemAddressReserve (total_aligned_size , addr_align , 0 , 0 )
394393 raise_if_driver_error (res )
395394 # Register undo for VA reservation
396- trans .append (
395+ trans .on_failure (
397396 lambda np = new_ptr , s = total_aligned_size : raise_if_driver_error (driver .cuMemAddressFree (np , s )[0 ])
398397 )
399398
400399 # Get the old allocation handle for remapping
401400 result , old_handle = driver .cuMemRetainAllocationHandle (buf .handle )
402401 raise_if_driver_error (result )
403- # Register undo for old_handle
404- trans .append (lambda h = old_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
402+ trans .on_exit (lambda h = old_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
405403
406404 # Unmap the old VA range (aligned previous size)
407405 aligned_prev_size = total_aligned_size - aligned_additional_size
@@ -417,28 +415,26 @@ def _remap_old() -> None:
417415 # TODO: consider logging this exception
418416 pass
419417
420- trans .append (_remap_old )
418+ trans .on_failure (_remap_old )
421419
422420 # Remap the old physical memory to the new VA range (aligned previous size)
423421 (res ,) = driver .cuMemMap (int (new_ptr ), aligned_prev_size , 0 , old_handle , 0 )
424422 raise_if_driver_error (res )
425423
426424 # Register undo for mapping
427- trans .append (lambda np = new_ptr , s = aligned_prev_size : raise_if_driver_error (driver .cuMemUnmap (np , s )[0 ]))
425+ trans .on_failure (lambda np = new_ptr , s = aligned_prev_size : raise_if_driver_error (driver .cuMemUnmap (np , s )[0 ]))
428426
429427 # Create new physical memory for the additional size
430428 res , new_handle = driver .cuMemCreate (aligned_additional_size , prop , 0 )
431429 raise_if_driver_error (res )
432-
433- # Register undo for new physical memory
434- trans .append (lambda h = new_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
430+ trans .on_exit (lambda h = new_handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
435431
436432 # Map the new physical memory to the extended portion (aligned offset)
437433 (res ,) = driver .cuMemMap (int (new_ptr ) + aligned_prev_size , aligned_additional_size , 0 , new_handle , 0 )
438434 raise_if_driver_error (res )
439435
440436 # Register undo for mapping
441- trans .append (
437+ trans .on_failure (
442438 lambda base = int (new_ptr ), offs = aligned_prev_size , s = aligned_additional_size : raise_if_driver_error (
443439 driver .cuMemUnmap (base + offs , s )[0 ]
444440 )
@@ -553,20 +549,20 @@ def allocate(self, size: int, *, stream: Stream | GraphBuilder | None = None) ->
553549 # ---- Create physical memory ----
554550 res , handle = driver .cuMemCreate (aligned_size , prop , 0 )
555551 raise_if_driver_error (res )
556- # Register undo for physical memory
557- trans .append (lambda h = handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
552+ # Drop the creation reference on either outcome; a successful mapping keeps the allocation alive.
553+ trans .on_exit (lambda h = handle : raise_if_driver_error (driver .cuMemRelease (h )[0 ]))
558554
559555 # ---- Reserve VA space ----
560556 # Potentially, use a separate size for the VA reservation from the physical allocation size
561557 res , ptr = driver .cuMemAddressReserve (aligned_size , addr_align , config .addr_hint , 0 )
562558 raise_if_driver_error (res )
563559 # Register undo for VA reservation
564- trans .append (lambda p = ptr , s = aligned_size : raise_if_driver_error (driver .cuMemAddressFree (p , s )[0 ]))
560+ trans .on_failure (lambda p = ptr , s = aligned_size : raise_if_driver_error (driver .cuMemAddressFree (p , s )[0 ]))
565561
566562 # ---- Map physical memory into VA ----
567563 (res ,) = driver .cuMemMap (ptr , aligned_size , 0 , handle , 0 )
568- trans .append (lambda p = ptr , s = aligned_size : raise_if_driver_error (driver .cuMemUnmap (p , s )[0 ]))
569564 raise_if_driver_error (res )
565+ trans .on_failure (lambda p = ptr , s = aligned_size : raise_if_driver_error (driver .cuMemUnmap (p , s )[0 ]))
570566
571567 # ---- Set access for owner + peers ----
572568 descs = self ._build_access_descriptors (prop )
@@ -600,14 +596,11 @@ def deallocate(self, ptr: DevicePointerType, size: int, *, stream: Stream | Grap
600596 from cuda .core ._stream import Stream_accept
601597
602598 Stream_accept (stream )
603- result , handle = driver .cuMemRetainAllocationHandle (ptr )
604- raise_if_driver_error (result )
599+ # The mapping owns the allocation; unmapping frees its backing memory when no external references remain.
605600 (result ,) = driver .cuMemUnmap (ptr , size )
606601 raise_if_driver_error (result )
607602 (result ,) = driver .cuMemAddressFree (ptr , size )
608603 raise_if_driver_error (result )
609- (result ,) = driver .cuMemRelease (handle )
610- raise_if_driver_error (result )
611604
612605 @property
613606 def is_device_accessible (self ) -> bool :
0 commit comments