Skip to content

Commit 2f0d2bf

Browse files
jeffkbkimfacebook-github-bot
authored andcommitted
Preserve heterogeneous inputs in ZORM micro-batches
Summary: Preserve valid heterogeneous model outputs and required inputs when batching ZORM updates. Jagged values retain variable cardinality, optional all-empty fields keep their semantics, and model-output normalization now runs consistently on synchronous and CPU-offloaded paths. Differential Revision: D115274546
1 parent ee4cc84 commit 2f0d2bf

4 files changed

Lines changed: 266 additions & 13 deletions

File tree

torchrec/metrics/cpu_offloaded_metric_module.py

Lines changed: 68 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -90,6 +90,7 @@ def log_event(*args: object, **kwargs: object) -> None:
9090
# an orphaned GLOO all_gather (a peer rank skipped the collective), an unbounded
9191
# join hangs the post_train_teardown lease and the whole job is StuckJob-killed.
9292
_COMPUTE_SHUTDOWN_JOIN_TIMEOUT_SEC: float = 300.0
93+
_JAGGED_LENGTHS_SUFFIX: str = "_lengths"
9394

9495

9596
def _format_thread_stack(thread: threading.Thread) -> str:
@@ -102,8 +103,22 @@ def _format_thread_stack(thread: threading.Thread) -> str:
102103
return "".join(traceback.format_stack(frame))
103104

104105

106+
def _jagged_value_keys(tensors: Mapping[str, torch.Tensor]) -> set[str]:
107+
return {
108+
key[: -len(_JAGGED_LENGTHS_SUFFIX)]
109+
for key in tensors
110+
if key.endswith(_JAGGED_LENGTHS_SUFFIX)
111+
and key[: -len(_JAGGED_LENGTHS_SUFFIX)] in tensors
112+
}
113+
114+
105115
def _job_batch_size(model_out: Dict[str, torch.Tensor]) -> int:
106-
sizes = [t.shape[0] for t in model_out.values() if t.dim() >= 1]
116+
jagged_value_keys = _jagged_value_keys(model_out)
117+
sizes = [
118+
tensor.shape[0]
119+
for key, tensor in model_out.items()
120+
if key not in jagged_value_keys and tensor.dim() >= 1
121+
]
107122
return max(sizes) if sizes else 1
108123

109124

@@ -131,11 +146,54 @@ def _merge_tensors_across_jobs(
131146
)
132147

133148

149+
def _concatenate_variable_cardinality(
150+
tensors: list[torch.Tensor],
151+
) -> torch.Tensor:
152+
rank = tensors[0].dim()
153+
trailing_shape = tensors[0].shape[1:]
154+
if any(
155+
tensor.dim() != rank or tensor.shape[1:] != trailing_shape
156+
for tensor in tensors[1:]
157+
):
158+
raise RecMetricException(
159+
"cannot concatenate tensors with shapes "
160+
f"{[tuple(tensor.shape) for tensor in tensors]}: dimensions after "
161+
"dim 0 must match"
162+
)
163+
return torch.cat(tensors, dim=0)
164+
165+
166+
def _merge_tensor_mappings(
167+
mappings: list[Mapping[str, torch.Tensor]],
168+
batch_sizes: list[int],
169+
mapping_name: str,
170+
) -> Dict[str, torch.Tensor]:
171+
first = mappings[0]
172+
jagged_value_keys = _jagged_value_keys(first)
173+
merged: Dict[str, torch.Tensor] = {}
174+
for key in first:
175+
tensors = [mapping[key] for mapping in mappings]
176+
try:
177+
if key in jagged_value_keys or all(
178+
tensor.numel() == 0 for tensor in tensors
179+
):
180+
merged[key] = _concatenate_variable_cardinality(tensors)
181+
else:
182+
merged[key] = _merge_tensors_across_jobs(tensors, batch_sizes)
183+
except (RecMetricException, RuntimeError) as error:
184+
raise RecMetricException(
185+
f"failed to merge {mapping_name} key {key!r}: {error}"
186+
) from error
187+
return merged
188+
189+
134190
def _merge_update_jobs(jobs: list[MetricUpdateJob]) -> MetricUpdateJob:
135191
"""
136192
Merge a list of MetricUpdateJobs into a single job by concatenating
137193
per-key tensors in `model_out` and any tensor entries in
138194
`kwargs['required_inputs']` along the batch dimension (dim=0).
195+
Values with a paired `<key>_lengths` tensor retain their jagged
196+
cardinality, and fields that are empty in every job remain empty.
139197
140198
Worker-side input batching: K mini-batches' worth of model outputs
141199
are passed as a single (concatenated) job, cutting the DtoH transfer,
@@ -152,20 +210,18 @@ def _merge_update_jobs(jobs: list[MetricUpdateJob]) -> MetricUpdateJob:
152210

153211
first = jobs[0]
154212
batch_sizes = [_job_batch_size(j.model_out) for j in jobs]
155-
merged_model_out: Dict[str, torch.Tensor] = {
156-
key: _merge_tensors_across_jobs([j.model_out[key] for j in jobs], batch_sizes)
157-
for key in first.model_out.keys()
158-
}
213+
merged_model_out = _merge_tensor_mappings(
214+
[job.model_out for job in jobs], batch_sizes, "model_out"
215+
)
159216

160217
merged_kwargs: Dict[str, Any] = dict(first.kwargs)
161218
required_inputs_first = first.kwargs.get("required_inputs")
162219
if isinstance(required_inputs_first, dict) and required_inputs_first:
163-
merged_required: Dict[str, torch.Tensor] = {
164-
key: _merge_tensors_across_jobs(
165-
[j.kwargs["required_inputs"][key] for j in jobs], batch_sizes
166-
)
167-
for key in required_inputs_first.keys()
168-
}
220+
merged_required = _merge_tensor_mappings(
221+
[job.kwargs["required_inputs"] for job in jobs],
222+
batch_sizes,
223+
"required_inputs",
224+
)
169225
merged_kwargs["required_inputs"] = merged_required
170226

171227
return MetricUpdateJob(
@@ -517,6 +573,7 @@ def _process_metric_update_job(self, metric_update_job: MetricUpdateJob) -> None
517573
if transfer_completed_event is not None:
518574
transfer_completed_event.synchronize()
519575

576+
cpu_model_out = self._prepare_model_out_for_metrics(cpu_model_out)
520577
labels, predictions, weights, required_inputs = parse_task_model_outputs(
521578
self.rec_tasks,
522579
cpu_model_out,

torchrec/metrics/metric_module.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -315,6 +315,17 @@ def load_state_dict_hook(
315315
f"Removed key '{key}' from state_dict for backward compatibility"
316316
)
317317

318+
def _prepare_model_out_for_metrics(
319+
self, model_out: Dict[str, torch.Tensor]
320+
) -> Dict[str, torch.Tensor]:
321+
"""Prepare model outputs before task and required-input parsing.
322+
323+
Subclasses can override this hook when their model-output format needs
324+
normalization. CPU-offloaded modules run it on the worker after transfer
325+
to CPU, keeping preparation off the trainer critical path.
326+
"""
327+
return model_out
328+
318329
def _update_rec_metrics(
319330
self, model_out: Dict[str, torch.Tensor], **kwargs: Any
320331
) -> None:
@@ -324,6 +335,7 @@ def _update_rec_metrics(
324335
"""
325336
# pyrefly: ignore[not-callable]
326337
if self.rec_metrics and self.rec_tasks:
338+
model_out = self._prepare_model_out_for_metrics(model_out)
327339
labels, predictions, weights, required_inputs = parse_task_model_outputs(
328340
self.rec_tasks, model_out, self.get_required_inputs()
329341
)

torchrec/metrics/tests/test_cpu_offloaded_metric_module.py

Lines changed: 141 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1921,7 +1921,147 @@ def test_unalignable_required_input_raises(self) -> None:
19211921
kwargs={"required_inputs": {"target_tensor": torch.tensor([7.0, 8.0])}},
19221922
),
19231923
]
1924-
with self.assertRaises(RecMetricException):
1924+
with self.assertRaisesRegex(
1925+
RecMetricException,
1926+
"failed to merge required_inputs key 'target_tensor'",
1927+
):
1928+
_merge_update_jobs(jobs)
1929+
1930+
def test_all_empty_unused_model_output_is_preserved(self) -> None:
1931+
jobs = [
1932+
MetricUpdateJob(
1933+
model_out={
1934+
"prediction": torch.tensor([0.1, 0.2]),
1935+
"optional_teacher_label": torch.empty(0),
1936+
},
1937+
kwargs={},
1938+
),
1939+
MetricUpdateJob(
1940+
model_out={
1941+
"prediction": torch.tensor([0.3, 0.4]),
1942+
"optional_teacher_label": torch.empty(0),
1943+
},
1944+
kwargs={},
1945+
),
1946+
]
1947+
1948+
merged = _merge_update_jobs(jobs)
1949+
1950+
self.assertEqual(merged.model_out["prediction"].shape, (4,))
1951+
self.assertEqual(merged.model_out["optional_teacher_label"].shape, (0,))
1952+
1953+
def test_all_empty_model_outputs_require_matching_trailing_shapes(self) -> None:
1954+
jobs = [
1955+
MetricUpdateJob(
1956+
model_out={
1957+
"prediction": torch.tensor([0.1, 0.2]),
1958+
"optional_teacher_label": torch.empty(0, 5),
1959+
},
1960+
kwargs={},
1961+
),
1962+
MetricUpdateJob(
1963+
model_out={
1964+
"prediction": torch.tensor([0.3, 0.4]),
1965+
"optional_teacher_label": torch.empty(0, 3),
1966+
},
1967+
kwargs={},
1968+
),
1969+
]
1970+
1971+
with self.assertRaisesRegex(
1972+
RecMetricException,
1973+
"failed to merge model_out key 'optional_teacher_label'.*"
1974+
"dimensions after dim 0 must match",
1975+
):
1976+
_merge_update_jobs(jobs)
1977+
1978+
def test_jagged_values_and_lengths_preserve_variable_cardinality(self) -> None:
1979+
jobs = [
1980+
MetricUpdateJob(
1981+
model_out={
1982+
"prediction": torch.tensor([0.1, 0.2]),
1983+
"campaign_id": torch.empty(0, dtype=torch.long),
1984+
"campaign_id_lengths": torch.tensor([0, 0]),
1985+
},
1986+
kwargs={},
1987+
),
1988+
MetricUpdateJob(
1989+
model_out={
1990+
"prediction": torch.tensor([0.3, 0.4]),
1991+
"campaign_id": torch.tensor([7, 7, 8]),
1992+
"campaign_id_lengths": torch.tensor([2, 1]),
1993+
},
1994+
kwargs={},
1995+
),
1996+
]
1997+
1998+
merged = _merge_update_jobs(jobs)
1999+
2000+
torch.testing.assert_close(
2001+
merged.model_out["prediction"], torch.tensor([0.1, 0.2, 0.3, 0.4])
2002+
)
2003+
torch.testing.assert_close(
2004+
merged.model_out["campaign_id"], torch.tensor([7, 7, 8])
2005+
)
2006+
torch.testing.assert_close(
2007+
merged.model_out["campaign_id_lengths"], torch.tensor([0, 0, 2, 1])
2008+
)
2009+
2010+
def test_jagged_required_inputs_preserve_variable_cardinality(self) -> None:
2011+
jobs = [
2012+
MetricUpdateJob(
2013+
model_out={"prediction": torch.tensor([0.1, 0.2])},
2014+
kwargs={
2015+
"required_inputs": {
2016+
"session_id": torch.empty(0, dtype=torch.long),
2017+
"session_id_lengths": torch.tensor([0, 0]),
2018+
}
2019+
},
2020+
),
2021+
MetricUpdateJob(
2022+
model_out={"prediction": torch.tensor([0.3, 0.4])},
2023+
kwargs={
2024+
"required_inputs": {
2025+
"session_id": torch.tensor([9, 9, 10]),
2026+
"session_id_lengths": torch.tensor([2, 1]),
2027+
}
2028+
},
2029+
),
2030+
]
2031+
2032+
merged = _merge_update_jobs(jobs)
2033+
2034+
torch.testing.assert_close(
2035+
merged.kwargs["required_inputs"]["session_id"],
2036+
torch.tensor([9, 9, 10]),
2037+
)
2038+
torch.testing.assert_close(
2039+
merged.kwargs["required_inputs"]["session_id_lengths"],
2040+
torch.tensor([0, 0, 2, 1]),
2041+
)
2042+
2043+
def test_unpaired_mixed_empty_tensor_remains_invalid(self) -> None:
2044+
jobs = [
2045+
MetricUpdateJob(
2046+
model_out={
2047+
"prediction": torch.tensor([0.1, 0.2]),
2048+
"ambiguous": torch.empty(0),
2049+
},
2050+
kwargs={},
2051+
),
2052+
MetricUpdateJob(
2053+
model_out={
2054+
"prediction": torch.tensor([0.3, 0.4]),
2055+
"ambiguous": torch.tensor([1.0]),
2056+
},
2057+
kwargs={},
2058+
),
2059+
]
2060+
2061+
with self.assertRaisesRegex(
2062+
RecMetricException,
2063+
"failed to merge model_out key 'ambiguous'",
2064+
):
19252065
_merge_update_jobs(jobs)
19262066

19272067
def test_scalar_model_out_expands_and_concats(self) -> None:

torchrec/metrics/tests/test_metric_module.py

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
import tempfile
1616
import unittest
1717
from collections.abc import Mapping
18-
from typing import Any, Callable, Dict, List, Optional
18+
from typing import Any, Callable, cast, Dict, List, Optional
1919
from unittest.mock import MagicMock, patch
2020

2121
import torch
@@ -101,6 +101,15 @@ def _update_rec_metrics(
101101
self.rec_metrics.update(predictions=predictions, labels=labels, weights=weights)
102102

103103

104+
class PreparingMetricModule(RecMetricModule):
105+
def _prepare_model_out_for_metrics(
106+
self, model_out: Dict[str, torch.Tensor]
107+
) -> Dict[str, torch.Tensor]:
108+
prepared_model_out = dict(model_out)
109+
prepared_model_out["task-prediction"] = prepared_model_out["raw-prediction"]
110+
return prepared_model_out
111+
112+
104113
class MetricModuleTest(unittest.TestCase):
105114
@seed_and_log
106115
def setUp(self) -> None:
@@ -189,6 +198,41 @@ def test_metric_module(self) -> None:
189198
self.assertTrue("optimizers-optimizers|learning_rate" in ret)
190199
dist.destroy_process_group()
191200

201+
def test_prepares_model_output_before_parsing(self) -> None:
202+
tasks = gen_test_tasks(["task"])
203+
rec_metric = MockRecMetric(
204+
world_size=1,
205+
my_rank=0,
206+
batch_size=2,
207+
tasks=tasks,
208+
)
209+
metric_module = PreparingMetricModule(
210+
batch_size=2,
211+
world_size=1,
212+
rec_tasks=tasks,
213+
rec_metrics=RecMetricList([rec_metric]),
214+
)
215+
raw_predictions = torch.tensor([0.25, 0.75])
216+
model_out = gen_test_batch(
217+
batch_size=2,
218+
label_name="task-label",
219+
prediction_name="unused-prediction",
220+
weight_name="task-weight",
221+
)
222+
model_out["raw-prediction"] = raw_predictions
223+
224+
metric_module.update(model_out)
225+
226+
self.assertEqual(1, rec_metric.update_called_count)
227+
actual_predictions = rec_metric.predictions_update_calls[0]
228+
self.assertIsInstance(actual_predictions, dict)
229+
self.assertTrue(
230+
torch.equal(
231+
raw_predictions,
232+
cast(Dict[str, torch.Tensor], actual_predictions)["task"],
233+
)
234+
)
235+
192236
def test_rectask_info(self) -> None:
193237
mock_optimizer = MockOptimizer()
194238
config = DefaultMetricsConfig

0 commit comments

Comments
 (0)