From 13fa120aef5faa0f04f2cb8c574a0dc1a533c6af Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 19:34:18 +0800 Subject: [PATCH] [bugfix] restore ODPS sessions by partition index on resume On resume, OdpsReader rebuilt the per-input session list as the sessions found in the checkpoint followed by the freshly created sessions from that count onwards. This assumed the checkpointed sessions were exactly the leading partitions in order, but a partition that produced no rows before the checkpoint (e.g. an empty one) never gets a key, and the key order follows the rank-major merge rather than partition order. A restored session could then land at the wrong position, leaving its partition's fresh session in the list and re-reading that partition from row 0. The checkpoint source key now carries the session's position ("{input_path}#{sess_idx}#{session_id}:{start}") and restore puts each validated session back at that index, so unconsumed partitions keep their own fresh session. Keys without an index come from an older version and raise a clear error. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MgZZz8Baa365GMyVuRFvZb --- tzrec/datasets/odps_dataset.py | 78 ++++++++------- tzrec/datasets/odps_dataset_test.py | 149 ++++++++++++++++++++++++++-- 2 files changed, 182 insertions(+), 45 deletions(-) diff --git a/tzrec/datasets/odps_dataset.py b/tzrec/datasets/odps_dataset.py index 294f2aa3..e2c1cb38 100644 --- a/tzrec/datasets/odps_dataset.py +++ b/tzrec/datasets/odps_dataset.py @@ -547,34 +547,34 @@ def _init_session(self) -> None: def _restore_sessions(self, checkpoint_state: Dict[str, int]) -> None: """Restore ODPS sessions from checkpoint state. - Parses session_ids from checkpoint keys and validates they are still active. - Raises RuntimeError if any session is expired/invalid. + Parses session ids and their positions from checkpoint keys, validates the + sessions are still active and puts each one back at its own position, so + partitions that never produced a key (e.g. empty ones) keep the freshly + created session at their own index. Keys of unknown input paths are skipped + with a warning. Raises RuntimeError if a key of a current input path has no + session index (written by an older version), if an index exceeds the current + number of partition sessions, or if any session is expired/invalid. - Checkpoint key format: "{input_path}#{session_id}:{start}" + Checkpoint key format: "{input_path}#{sess_idx}#{session_id}:{start}" Args: checkpoint_state: Checkpoint state dict mapping source_key to row index. """ - # Parse unique session_ids from checkpoint keys - session_ids_by_input: Dict[str, List[str]] = {} # {input_path: [session_ids]} + sess_by_input: Dict[str, Dict[int, str]] = {} # {input_path: {idx: sess_id}} for key in checkpoint_state.keys(): - # Parse: "{input_path}#{session_id}:{start}" - last_colon = key.rfind(":") - if last_colon == -1: + prefix, sep, _ = key.rpartition(":") + if not sep: continue - prefix = key[:last_colon] # "{input_path}#{session_id}" - hash_idx = prefix.rfind("#") - if hash_idx == -1: + input_and_idx, sep, session_id = prefix.rpartition("#") + if not sep: continue - input_path = prefix[:hash_idx] - session_id = prefix[hash_idx + 1 :] - if input_path not in session_ids_by_input: - session_ids_by_input[input_path] = [] - if session_id not in session_ids_by_input[input_path]: - session_ids_by_input[input_path].append(session_id) - - # Restore sessions for each input_path - for input_path, session_ids in session_ids_by_input.items(): + input_path, sep, sess_idx = input_and_idx.rpartition("#") + if not sep or not sess_idx.isdigit(): + # "{input_path}#{session_id}" key of an older version + input_path, sess_idx = input_and_idx, "-1" + sess_by_input.setdefault(input_path, {})[int(sess_idx)] = session_id + + for input_path, idx_to_sess in sess_by_input.items(): if input_path not in self._input_to_sess: logger.warning( f"Checkpoint contains unknown input_path: {input_path}. " @@ -583,9 +583,23 @@ def _restore_sessions(self, checkpoint_state: Dict[str, int]) -> None: continue _, table_name, _, _ = _parse_table_path(input_path) client = self._table_to_cli[table_name] + sess_reqs = list(self._input_to_sess[input_path]) - restored_sess_reqs = [] - for session_id in session_ids: + for sess_idx, session_id in idx_to_sess.items(): + if sess_idx < 0: + raise RuntimeError( + f"Cannot resume from checkpoint: ODPS session {session_id} " + f"for {input_path} has no session index, the checkpoint was " + "written by an older TorchEasyRec version. " + "Please restart training from scratch." + ) + if sess_idx >= len(sess_reqs): + raise RuntimeError( + f"Cannot resume from checkpoint: ODPS session {session_id} " + f"for {input_path} has index {sess_idx} but only " + f"{len(sess_reqs)} partition sessions exist now. " + "Please restart training from scratch." + ) sess_req = SessionRequest(session_id=session_id) try: resp = client.get_read_session(sess_req) @@ -595,7 +609,7 @@ def _restore_sessions(self, checkpoint_state: Dict[str, int]) -> None: f"for {input_path} has expired. Row order may have changed." " Please restart training from scratch." ) - restored_sess_reqs.append(sess_req) + sess_reqs[sess_idx] = sess_req self._sess_row_counts[session_id] = resp.record_count except ODPSError as e: raise RuntimeError( @@ -603,19 +617,7 @@ def _restore_sessions(self, checkpoint_state: Dict[str, int]) -> None: f"for {input_path} is invalid. Error: {e}. " "Please restart training from scratch." ) from e - - # When orderby partition, partitions are consumed in order. - # Checkpoint only contains consumed partition sessions. - # Merge: restored sessions + unconsumed partition sessions (newly created) - current_sessions = self._input_to_sess.get(input_path, []) - n_restored = len(restored_sess_reqs) - if n_restored < len(current_sessions): - # Keep sessions for unconsumed partitions - self._input_to_sess[input_path] = ( - restored_sess_reqs + current_sessions[n_restored:] - ) - else: - self._input_to_sess[input_path] = restored_sess_reqs + self._input_to_sess[input_path] = sess_reqs def load_state_dict(self, state: Optional[Dict[str, int]]) -> None: """Set checkpoint state and restore sessions. @@ -636,10 +638,10 @@ def to_batches( for input_path in self._input_path.split(","): _, table_name, _, _ = _parse_table_path(input_path) client = self._table_to_cli[table_name] - for sess_req in self._input_to_sess[input_path]: + for sess_idx, sess_req in enumerate(self._input_to_sess[input_path]): sources.append( ( - f"{input_path}#{sess_req.session_id}", + f"{input_path}#{sess_idx}#{sess_req.session_id}", self._sess_row_counts[sess_req.session_id], client, sess_req, diff --git a/tzrec/datasets/odps_dataset_test.py b/tzrec/datasets/odps_dataset_test.py index 2ce50dca..a56538c7 100644 --- a/tzrec/datasets/odps_dataset_test.py +++ b/tzrec/datasets/odps_dataset_test.py @@ -21,6 +21,7 @@ import pyarrow as pa import requests from odps import ODPS +from odps.apis.storage_api import SessionRequest, SessionStatus from odps.errors import ODPSError from parameterized import parameterized from torch import distributed as dist @@ -255,7 +256,7 @@ def test_odps_dataset_checkpoint_metadata(self): self.assertIsNotNone(batch.checkpoint_info) self.assertIsInstance(batch.checkpoint_info, dict) - # Checkpoint keys should be in format "{input_path}#{session_id}:{start}" + # Checkpoint keys are "{input_path}#{sess_idx}#{session_id}:{start}" for key, value in batch.checkpoint_info.items(): self.assertIn(":", key) self.assertIn("#", key) @@ -263,11 +264,9 @@ def test_odps_dataset_checkpoint_metadata(self): self.assertEqual(len(parts), 2) # start should be numeric self.assertTrue(parts[1].isdigit()) - # Verify session_id is present (between # and last :) - prefix = parts[0] - hash_idx = prefix.rfind("#") - self.assertGreater(hash_idx, 0) - session_id = prefix[hash_idx + 1 :] + input_path, sess_idx, session_id = parts[0].split("#") + self.assertGreater(len(input_path), 0) + self.assertTrue(sess_idx.isdigit()) self.assertGreater(len(session_id), 0) # Value should be a non-negative integer self.assertIsInstance(value, int) @@ -438,9 +437,86 @@ def test_odps_dataset_checkpoint_resume_orderby_partition(self): ) self.assertEqual( restored_sess_list[0].session_id, - list(first_checkpoint_state.keys())[0].split("#")[1].split(":")[0], + list(first_checkpoint_state.keys())[0].rsplit("#", 1)[1].split(":")[0], ) + @unittest.skipIf( + os.environ.get("ODPS_CONFIG_FILE_PATH", "") == "" + and os.environ.get("ALIBABA_CLOUD_ECS_METADATA", "") == "", + "odps config not found", + ) + def test_odps_dataset_checkpoint_resume_empty_partition(self): + """Test checkpoint resume keeps sessions at their partition position.""" + account, odps_endpoint = _create_odps_account() + project = self.test_project + self.o = ODPS(account=account, project=project, endpoint=odps_endpoint) + feature_cfgs = self._create_test_table_and_feature_cfgs() + features = create_features(feature_cfgs, fg_mode=FgMode.FG_DAG) + table_name = f"test_odps_dataset_{self.test_suffix}" + # an empty partition between two 10000-row partitions never yields a + # checkpoint key, so the consumed sessions are not a prefix of the list + self.o.get_table(table_name).create_partition("dt=20240318") + input_path = ( + f"odps://{project}/tables/{table_name}/dt=20240319&dt=20240318&dt=20240320" + ) + data_config = data_pb2.DataConfig( + batch_size=1024, + dataset_type=data_pb2.DatasetType.OdpsDataset, + fg_mode=FgMode.FG_DAG, + label_fields=["label"], + is_orderby_partition=True, + odps_data_quota_name=self.test_quota, + ) + + dataset1 = OdpsDataset( + data_config=data_config, features=features, input_path=input_path + ) + self.assertEqual(len(list(dataset1._reader._input_to_sess.values())[0]), 3) + dataloader1 = DataLoader( + dataset=dataset1, + batch_size=None, + num_workers=2, + pin_memory=True, + collate_fn=lambda x: x, + ) + iterator1 = iter(dataloader1) + + # consume into the last partition + checkpoint_state = {} + num_consumed = 0 + for _ in range(20): + batch = next(iterator1) + update_dataloder_state(checkpoint_state, batch.checkpoint_info.copy()) + num_consumed += len(batch.labels["label"]) + if any("#2#" in key for key in checkpoint_state): + break + self.assertTrue(any("#2#" in key for key in checkpoint_state)) + self.assertFalse(any("#1#" in key for key in checkpoint_state)) + del iterator1 + del dataloader1 + + dataset2 = OdpsDataset( + data_config=data_config, features=features, input_path=input_path + ) + fresh_sess_list = list(dataset2._reader._input_to_sess.values())[0] + dataset2.load_state_dict(checkpoint_state) + + ckpt_sess_ids = { + int(key.rsplit("#", 2)[1]): key.rsplit("#", 1)[1].split(":")[0] + for key in checkpoint_state + } + restored_sess_list = list(dataset2._reader._input_to_sess.values())[0] + self.assertEqual(len(restored_sess_list), 3) + self.assertEqual(restored_sess_list[0].session_id, ckpt_sess_ids[0]) + self.assertEqual( + restored_sess_list[1].session_id, fresh_sess_list[1].session_id + ) + self.assertEqual(restored_sess_list[2].session_id, ckpt_sess_ids[2]) + + # the remaining rows are read exactly once + num_remaining = sum(len(batch.labels["label"]) for batch in dataset2) + self.assertEqual(num_consumed + num_remaining, 20000) + def _test_odps_dataset_with_sampler(self, id_type="bigint", schema=None): account, odps_endpoint = _create_odps_account() project = self.test_project @@ -734,6 +810,65 @@ def read_rows_arrow(self, read_req): return self.sentinel +class _StubScanClient: + def __init__(self, status=SessionStatus.NORMAL): + self.status = status + + def get_read_session(self, sess_req): + return SimpleNamespace(session_status=self.status, record_count=100) + + +class OdpsRestoreSessionsTest(unittest.TestCase): + input_path = "odps://p/tables/t/dt=1&dt=2&dt=3" + + def _reader(self): + reader = odps_dataset.OdpsReader.__new__(odps_dataset.OdpsReader) + reader._input_to_sess = { + self.input_path: [SessionRequest(session_id=f"new{i}") for i in range(3)] + } + reader._table_to_cli = {"t": _StubScanClient()} + reader._sess_row_counts = {"new0": 100, "new1": 0, "new2": 100} + return reader + + def test_restore_keeps_partition_position(self): + reader = self._reader() + reader._restore_sessions( + { + f"{self.input_path}#2#old2:0": 39, + f"{self.input_path}#0#old0:50": 99, + f"{self.input_path}#0#old0:0": 49, + "__data_ts_watermark__": 5, + } + ) + self.assertEqual( + [x.session_id for x in reader._input_to_sess[self.input_path]], + ["old0", "new1", "old2"], + ) + self.assertEqual(reader._sess_row_counts["old2"], 100) + + def test_restore_skips_unknown_input_path(self): + reader = self._reader() + with mock.patch.object(odps_dataset, "logger") as m_logger: + reader._restore_sessions( + {"odps://p/tables/other/dt=1#0#old0:0": 9, "/data/a#b.parquet:0": 9} + ) + self.assertEqual(m_logger.warning.call_count, 2) + self.assertEqual( + [x.session_id for x in reader._input_to_sess[self.input_path]], + ["new0", "new1", "new2"], + ) + + def test_restore_rejects_key_without_session_index(self): + reader = self._reader() + with self.assertRaisesRegex(RuntimeError, "older TorchEasyRec version"): + reader._restore_sessions({f"{self.input_path}#old0:0": 9}) + + def test_restore_rejects_session_index_out_of_range(self): + reader = self._reader() + with self.assertRaisesRegex(RuntimeError, "only 3 partition sessions"): + reader._restore_sessions({f"{self.input_path}#3#old3:0": 9}) + + class OdpsStorageErrorLogTest(unittest.TestCase): def test_read_rows_arrow_logs_session_and_reraises(self): client = _FailingStorageClient()