11import math
2+ import warnings
23from collections .abc import Callable , Sequence
34from functools import partial
45
@@ -77,18 +78,18 @@ def _differentiate(self, jac_outputs: Sequence[Tensor], /) -> tuple[Tensor, ...]
7778 jacs_chunks : list [tuple [Tensor , ...]] = []
7879
7980 # First differentiations: always retain graph
80- get_vjp_retain = partial (self ._get_vjp , retain_graph = True )
8181 for i in range (n_chunks - 1 ):
8282 start = i * max_chunk_size
8383 end = (i + 1 ) * max_chunk_size
8484 jac_outputs_chunk = [jac_output [start :end ] for jac_output in jac_outputs ]
85- jacs_chunks .append (_get_jacs_chunk (jac_outputs_chunk , get_vjp_retain ))
85+ jacs_chunks .append (_get_jacs_chunk (jac_outputs_chunk , self . _get_vjp , retain_graph = True ))
8686
8787 # Last differentiation: retain the graph only if self.retain_graph==True
88- get_vjp_last = partial (self ._get_vjp , retain_graph = self .retain_graph )
8988 start = (n_chunks - 1 ) * max_chunk_size
9089 jac_outputs_chunk = [jac_output [start :] for jac_output in jac_outputs ]
91- jacs_chunks .append (_get_jacs_chunk (jac_outputs_chunk , get_vjp_last ))
90+ jacs_chunks .append (
91+ _get_jacs_chunk (jac_outputs_chunk , self ._get_vjp , retain_graph = self .retain_graph )
92+ )
9293
9394 n_inputs = len (self .inputs )
9495 if len (jacs_chunks ) == 1 :
@@ -102,7 +103,8 @@ def _differentiate(self, jac_outputs: Sequence[Tensor], /) -> tuple[Tensor, ...]
102103
103104def _get_jacs_chunk (
104105 jac_outputs_chunk : list [Tensor ],
105- get_vjp : Callable [[Sequence [Tensor ]], tuple [Tensor , ...]],
106+ get_vjp : Callable ,
107+ retain_graph : bool ,
106108) -> tuple [Tensor , ...]:
107109 """
108110 Computes the jacobian matrix chunk corresponding to the provided get_vjp function, either by
@@ -112,8 +114,28 @@ def _get_jacs_chunk(
112114 """
113115
114116 chunk_size = jac_outputs_chunk [0 ].shape [0 ]
117+
118+ def _vmap_target (grad_outputs : Sequence [Tensor ]) -> tuple [Tensor , ...]:
119+ return get_vjp (grad_outputs , retain_graph = retain_graph )
120+
115121 if chunk_size == 1 :
116122 grad_outputs = [tensor .squeeze (0 ) for tensor in jac_outputs_chunk ]
117- gradients = get_vjp (grad_outputs )
123+ gradients = _vmap_target (grad_outputs )
118124 return tuple (gradient .unsqueeze (0 ) for gradient in gradients )
119- return torch .vmap (get_vjp , chunk_size = chunk_size )(jac_outputs_chunk )
125+
126+ try :
127+ return torch .vmap (_vmap_target , chunk_size = chunk_size )(jac_outputs_chunk )
128+ except RuntimeError as e :
129+ warnings .warn (
130+ f"torch.vmap failed with RuntimeError: { e } . "
131+ "Falling back to sequential differentiation. To suppress this warning, "
132+ "explicitly provide `chunk_size=1` or `parallel_chunk_size=1`."
133+ )
134+ # Fallback to sequential execution if vmap fails (e.g., due to AMP mixed precision issues)
135+ jacs = []
136+ for i in range (chunk_size ):
137+ grad_outputs = [tensor [i ] for tensor in jac_outputs_chunk ]
138+ # Retain graph for all elements except the last one (if retain_graph is False)
139+ should_retain = retain_graph or (i < chunk_size - 1 )
140+ jacs .append (get_vjp (grad_outputs , retain_graph = should_retain ))
141+ return tuple (torch .stack (grads ) for grads in zip (* jacs , strict = True ))
0 commit comments