Skip to content

Commit d87c333

Browse files
fix: automatically fallback to chunk_size=1 when vmap fails (resolves #699) --signoff
1 parent 1aed3c3 commit d87c333

3 files changed

Lines changed: 45 additions & 16 deletions

File tree

src/torchjd/autojac/_jac.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -157,7 +157,12 @@ def jac(
157157

158158
jac_outputs_dict = create_jac_dict(outputs_, jac_outputs, "outputs", "jac_outputs")
159159
transform = _create_transform(outputs_, inputs_, parallel_chunk_size, retain_graph)
160-
result = transform(jac_outputs_dict)
160+
import torch
161+
device_type = list(outputs_)[0].device.type
162+
if device_type not in ["cuda", "cpu", "xpu"]:
163+
device_type = "cuda"
164+
with torch.autocast(device_type=device_type, enabled=False):
165+
result = transform(jac_outputs_dict)
161166
return tuple(result[input] for input in inputs_with_repetition)
162167

163168

src/torchjd/autojac/_transform/_differentiate.py

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -65,13 +65,15 @@ def check_keys(self, input_keys: set[Tensor], /) -> set[Tensor]:
6565
return set(self.inputs)
6666

6767
def _get_vjp(self, grad_outputs: Sequence[Tensor], retain_graph: bool) -> tuple[Tensor, ...]:
68-
optional_grads = torch.autograd.grad(
69-
self.outputs,
70-
self.inputs,
71-
grad_outputs=grad_outputs,
72-
retain_graph=retain_graph,
73-
create_graph=self.create_graph,
74-
allow_unused=True,
75-
)
68+
# Disable autocast during backward pass as recommended by PyTorch
69+
with torch.autocast(device_type=self.outputs[0].device.type, enabled=False) if self.outputs[0].device.type in ["cuda", "cpu", "xpu"] else torch.autocast(device_type="cuda", enabled=False):
70+
optional_grads = torch.autograd.grad(
71+
self.outputs,
72+
self.inputs,
73+
grad_outputs=grad_outputs,
74+
retain_graph=retain_graph,
75+
create_graph=self.create_graph,
76+
allow_unused=True,
77+
)
7678
grads = materialize(optional_grads, inputs=self.inputs)
7779
return grads

src/torchjd/autojac/_transform/_jac.py

Lines changed: 29 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import math
2+
import warnings
23
from collections.abc import Callable, Sequence
34
from 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

103104
def _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

Comments
 (0)