Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 40 additions & 38 deletions tzrec/datasets/odps_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Low-med: this raise fires for any key shaped ...#non-digit:..., including keys this reader never produced, which pre-diff code warned-and-skipped. Two concrete paths:

  • main.py:839-840 also loads dataloader_state.json from an external fine_tune_checkpoint source. An older-tzrec source checkpoint whose ODPS keys reference unrelated tables (fine-tuning onto different inputs) previously fell into the unknown-input_path warning below and was skipped; now it aborts startup on every rank with a misleading "restart from scratch".
  • Parquet source ids are {file_path}:{start} (parquet_dataset.py:276) and # is legal in POSIX/OSS paths, so e.g. /data/train#v2/part-0.parquet:1024 in a mixed-history state file hard-fails an ODPS job.

Suggestion: scope the raise to keys that claim one of this job's own inputs — for an old-format key, input_and_idx is the old input path, so raise only when it (or the parsed input_path) is in self._input_to_sess, otherwise fall through to the existing unknown-path warning. That keeps the intended hard failure for genuinely old checkpoints of the resumed table while restoring the tolerant skip for foreign keys.

# "{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}. "
Expand All @@ -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)
Expand All @@ -595,27 +609,15 @@ 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(
f"Cannot resume from checkpoint: ODPS session {session_id} "
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.
Expand All @@ -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,
Expand Down
149 changes: 142 additions & 7 deletions tzrec/datasets/odps_dataset_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -255,19 +256,17 @@ 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}"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The comment now claims the three-field format, but none of the assertions below actually distinguish it — they all pass identically under the old {input_path}#{session_id}:{start} layout, so a regression that drops sess_idx from to_batches would only be caught by the credential-gated resume tests. Cheap pin, since input paths contain no #: segments = prefix.split("#") → assertEqual(len(segments), 3) and assertTrue(segments[1].isdigit()).

for key, value in batch.checkpoint_info.items():
self.assertIn(":", key)
self.assertIn("#", key)
parts = key.rsplit(":", 1)
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)
Expand Down Expand Up @@ -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):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The two new RuntimeError branches — old-format key (odps_dataset.py:571) and out-of-range index (:591) — are the user-facing guardrails of this fix but have no committed coverage; every test that exercises _restore_sessions is gated behind ODPS credentials, and even those never hit the error paths.

Both are cheap to test offline in the style of OdpsStorageErrorLogTest below. The old-format raise fires during key parsing before any attribute access, so it works with a bare stub:

with self.assertRaisesRegex(RuntimeError, "older TorchEasyRec"):
    odps_dataset.OdpsReader._restore_sessions(
        SimpleNamespace(), {"odps://p/tables/t/dt=x#sessid:10": 0}
    )

The bounds check only needs _input_to_sess / _table_to_cli stubs (a SimpleNamespace with a 1-element session list and a key claiming index 5). Worth adding so the guards don't silently regress in non-credentialed CI lanes.

"""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
Expand Down Expand Up @@ -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()
Expand Down
Loading