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()