diff --git a/requirements/runtime.txt b/requirements/runtime.txt index c43dadbd..b23a4a5e 100644 --- a/requirements/runtime.txt +++ b/requirements/runtime.txt @@ -3,6 +3,7 @@ anytree common_io @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/common_io-0.4.1%2Btunnel-py2.py3-none-any.whl confluent-kafka fbgemm-gpu==1.7.0 +feature_store_py @ https://feature-store-py.oss-cn-beijing.aliyuncs.com/package/feature_store_py-2.2.7-py3-none-any.whl fsspec graphlearn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/graphlearn/graphlearn-1.3.8-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" graphlearn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/graphlearn/graphlearn-1.3.8-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" diff --git a/tzrec/main.py b/tzrec/main.py index 2a9b631c..6150f366 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -606,11 +606,9 @@ def run_eval(step: int, epoch: int) -> None: _model.on_train_end() if delta_embedding_dumper is not None: # Flush the trailing partial interval before the final checkpoint. - # final_dump skips dump-boundary steps already written by maybe_dump, - # so it never overwrites their shards with an empty file. Ranks can - # reach here at different i_step (independent dataloader exhaustion with - # check_all_workers_data_status=False), so final_dump all-reduces the - # step across ranks to keep one complete shard set per step dir. + # final_dump skips dump-boundary steps already written by maybe_dump + # (all ranks run the same step count, so every rank participated in + # those dumps and reaches the same final step). delta_embedding_dumper.final_dump(i_step) _log_train( @@ -900,6 +898,8 @@ def train_and_evaluate( with open(os.path.join(pipeline_config.model_dir, "version"), "w") as f: f.write(tzrec_version + "\n") + if delta_embedding_dumper is not None: + delta_embedding_dumper.start() # when slice batch by sample cost, data on all workers may not be balanced check_all_workers_data_status = data_config.HasField("batch_cost_size") _train_and_evaluate( @@ -922,6 +922,12 @@ def train_and_evaluate( dense_ema=dense_ema, export_config=pipeline_config.export_config, ) + # Drain background uploads only after training succeeds. A training failure + # terminates the whole job (torchrun tears down every rank) and pending + # in-memory deltas are intentionally abandoned: the restarted run re-dumps + # from the latest checkpoint, so there is nothing to roll back or undo. + if delta_embedding_dumper is not None: + delta_embedding_dumper.close() if is_local_rank_zero: logger.info("Train and Evaluate Finished.") diff --git a/tzrec/protos/train.proto b/tzrec/protos/train.proto index 05210303..d254199e 100644 --- a/tzrec/protos/train.proto +++ b/tzrec/protos/train.proto @@ -29,16 +29,71 @@ message GradClipping { optional bool enable_global_grad_clip = 4 [default = false]; } +message FeatureStoreConfig { + // Cloud credentials (AK/SK/STS) are resolved at runtime through the + // Alibaba Cloud default credential provider chain (alibabacloud_credentials). + // FeatureDB credentials are read from FEATUREDB_USERNAME/FEATUREDB_PASSWORD. + // FeatureStore control-plane region. An explicitly empty value falls back + // to ALIBABA_CLOUD_REGION at runtime. + required string region = 1; + // Existing FeatureStore project name and target DynamicEmbedding FeatureView + // name. The view is validated at startup and created when it does not exist, + // together with a default FeatureStore entity that is auto-created on demand. + required string project_name = 2; + required string feature_view_name = 3; + // FeatureDB version for this incremental training run. + required string version = 5; + + + // Optional FeatureStore control-plane endpoint override, passed directly to + // the FeatureStore SDK, which owns endpoint handling and connection errors. + optional string endpoint = 6; + // Maximum records submitted before each SDK write_flush() completion gate. + optional uint32 upload_batch_size = 7 [default = 1000]; + // Total attempts per rank for one step. All retries reuse the version; + // each attempt reserves a newer monotonic ts range and fully replays the + // rank's delta so incremental readers cannot miss a partial earlier attempt. + optional uint32 max_retries = 8 [default = 3]; + optional uint32 retry_backoff_secs = 9 [default = 5]; + // Maximum time normal training shutdown waits for the uploader to drain. + optional uint32 shutdown_timeout_secs = 10 [default = 600]; + // Apply back-pressure when too many completed dumps await upload; bounds + // the per-rank memory retained by pending in-memory deltas. + optional uint32 max_pending_steps = 11 [default = 32]; + optional uint32 poll_interval_secs = 12 [default = 1]; + // Creation settings used only when feature_view_name does not yet exist. + optional uint32 feature_view_ttl_secs = 13 [default = 1296000]; + optional uint32 feature_view_shard_count = 14 [default = 20]; + optional uint32 feature_view_replication_count = 15 [default = 1]; + reserved 4, 16; + reserved "feature_entity_name", "allow_custom_endpoint"; + // Also write the local per-rank delta parquet files while uploading to + // FeatureStore (e.g. for the offline readback checker); by default the + // delta is handed to the uploader in memory and never touches disk. + optional bool retain_local_dump = 17 [default = false]; + // Wire format for delta upload. "ARROW" (default) streams a columnar Arrow + // IPC batch through write_features_arrow(), avoiding the JSON path's per-row + // dict construction and embedding deep-copy; "JSON" keeps the legacy + // write_features() per-row payload. Both paths use MERGE write_mode. + optional string upload_format = 18 [default = "ARROW"]; +} + message DeltaEmbeddingDumpConfig { // MC/ZCH features are not supported; use dynamicemb for delta dump. - // dump touched ids and their latest embedding every N training steps. Larger - // intervals retain a longer id window in memory; auto compaction reduces + // Dump touched ids and their latest embedding every N training steps. Do not + // set this together with dump_interval_minutes. + // Larger intervals retain a longer id window in memory; auto compaction reduces // per-batch tensor buildup but unique ids still scale with the interval. optional uint32 dump_interval_steps = 1 [default = 1000]; // output directory. default is ${model_dir}/delta_embedding_dump optional string output_dir = 2; // parquet file prefix optional string file_prefix = 3 [default = "delta_embedding"]; + // Presence enables best-effort per-rank background upload to FeatureStore. + optional FeatureStoreConfig feature_store_config = 4; + // Dump after this many elapsed minutes. The timer starts when training starts. + // Do not set this together with dump_interval_steps. + optional uint32 dump_interval_minutes = 5; } message TrainConfig { diff --git a/tzrec/tools/feature_store/__init__.py b/tzrec/tools/feature_store/__init__.py new file mode 100644 index 00000000..eedc773b --- /dev/null +++ b/tzrec/tools/feature_store/__init__.py @@ -0,0 +1,10 @@ +# Copyright (c) 2026, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/tzrec/tools/feature_store/check_feature_store_delta.py b/tzrec/tools/feature_store/check_feature_store_delta.py new file mode 100644 index 00000000..541d2275 --- /dev/null +++ b/tzrec/tools/feature_store/check_feature_store_delta.py @@ -0,0 +1,535 @@ +# Copyright (c) 2026, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +r"""Read back sampled delta embeddings from FeatureStore. + +The tool scans the local delta parquet output for the latest (or specified) +step, samples keys from its shard set, and queries the configured explicit +FeatureDB version through ``DynamicEmbeddingFeatureView.get_online_features``. + +Local parquet files are only produced alongside FeatureStore uploads when +``feature_store_config.retain_local_dump`` is enabled. + +Example:: + + export ALIBABA_CLOUD_ACCESS_KEY_ID=YOUR_ACCESS_KEY_ID + export ALIBABA_CLOUD_ACCESS_KEY_SECRET=YOUR_ACCESS_KEY_SECRET + export FEATUREDB_USERNAME=YOUR_FEATUREDB_USERNAME + export FEATUREDB_PASSWORD=YOUR_FEATUREDB_PASSWORD + python -m tzrec.tools.feature_store.check_feature_store_delta \ + --pipeline_config path/to/pipeline.config \ + --output_dir path/to/delta_embedding_dump \ + --sample_count 10 +""" + +import argparse +import glob +import inspect +import json +import os +import sys +from collections import defaultdict +from dataclasses import dataclass, field +from typing import Any, DefaultDict, Dict, List, Optional, Sequence, Tuple + +import numpy as np +import numpy.typing as npt +import pyarrow.parquet as pq + +from tzrec.utils import config_util +from tzrec.utils.feature_store_delta_uploader import ( + FEATURE_STORE_PK_FIELD, + FEATURE_STORE_SK_FIELD, + FEATURE_STORE_VALUE_FIELD, + FeatureStoreUploadSettings, +) + + +@dataclass(frozen=True) +class LocalSample: + """One local parquet record selected for remote readback.""" + + embedding_name: str + key_id: int + embedding: npt.NDArray[np.float32] = field(repr=False) + source_path: str + + +def resolve_output_dir( + pipeline_config_path: str, + model_dir: str, + configured_output_dir: str, + output_dir_override: Optional[str], +) -> str: + """Resolve the local delta outbox, including relocated pipeline configs. + + Args: + pipeline_config_path: Source pipeline config path. + model_dir: Model directory from the pipeline config. + configured_output_dir: Explicit delta dump output directory, if any. + output_dir_override: Command-line output directory override, if any. + + Returns: + Absolute delta outbox directory path. + """ + if output_dir_override: + return os.path.abspath(output_dir_override) + if configured_output_dir: + return os.path.abspath(configured_output_dir) + + configured_path = os.path.abspath(os.path.join(model_dir, "delta_embedding_dump")) + if os.path.isdir(configured_path): + return configured_path + + colocated_path = os.path.join( + os.path.dirname(os.path.abspath(pipeline_config_path)), + "delta_embedding_dump", + ) + if os.path.isdir(colocated_path): + return colocated_path + return configured_path + + +def resolve_upload_step( + output_dir: str, + file_prefix: str, + world_size: int, + global_step: Optional[int] = None, +) -> Tuple[int, List[str]]: + """Find the latest (or specified) step with complete parquet shards. + + Args: + output_dir: Delta parquet outbox directory. + file_prefix: Scoped file prefix for parquet filenames. + world_size: Expected number of rank shards per step. + global_step: Specific step to inspect, or None for the latest. + + Returns: + Tuple of (global_step, shard_paths). + + Raises: + FileNotFoundError: If no complete shard set is found. + """ + if global_step is not None: + if global_step <= 0: + raise ValueError("global_step must be > 0") + paths = _shard_paths_for_step(output_dir, file_prefix, global_step, world_size) + if not paths: + raise FileNotFoundError( + f"no delta parquet shards found for step {global_step} " + f"under {output_dir}" + ) + return global_step, paths + + best_step = -1 + best_paths: List[str] = [] + if world_size == 1: + pattern = os.path.join(output_dir, f"{file_prefix}_step_*.parquet") + for path in glob.glob(pattern): + basename = os.path.basename(path) + step_str = basename.replace(f"{file_prefix}_step_", "").replace( + ".parquet", "" + ) + try: + step = int(step_str) + except ValueError: + continue + if step > best_step: + best_step = step + best_paths = [path] + else: + pattern = os.path.join(output_dir, "step_*") + for step_dir in sorted(glob.glob(pattern)): + if not os.path.isdir(step_dir): + continue + dir_name = os.path.basename(step_dir) + step_str = dir_name.replace("step_", "") + try: + step = int(step_str) + except ValueError: + continue + paths = _shard_paths_for_step(output_dir, file_prefix, step, world_size) + if paths and step > best_step: + best_step = step + best_paths = paths + + if best_step <= 0: + raise FileNotFoundError( + f"no complete delta parquet shard set found under {output_dir}" + ) + return best_step, best_paths + + +def _shard_paths_for_step( + output_dir: str, file_prefix: str, global_step: int, world_size: int +) -> List[str]: + """Resolve expected shard paths for one step, returning them if complete.""" + if world_size == 1: + path = os.path.join(output_dir, f"{file_prefix}_step_{global_step}.parquet") + return [path] if os.path.isfile(path) else [] + step_dir = os.path.join(output_dir, f"step_{global_step}") + paths = [ + os.path.join( + step_dir, + f"{file_prefix}_step_{global_step}_rank_{rank}_of_{world_size}.parquet", + ) + for rank in range(world_size) + ] + if all(os.path.isfile(path) for path in paths): + return paths + return [] + + +def sample_local_records( + parquet_paths: Sequence[str], + sample_count: int, + embedding_name: Optional[str] = None, +) -> List[LocalSample]: + """Read a bounded set of unique records from canonical parquet shards.""" + if sample_count <= 0: + raise ValueError("sample_count must be > 0") + + columns = [ + FEATURE_STORE_PK_FIELD, + FEATURE_STORE_SK_FIELD, + FEATURE_STORE_VALUE_FIELD, + ] + samples: List[LocalSample] = [] + seen: set[Tuple[str, int]] = set() + for path in parquet_paths: + parquet_file = pq.ParquetFile(path) + missing_columns = [ + name for name in columns if name not in parquet_file.schema_arrow.names + ] + if missing_columns: + raise ValueError( + f"delta parquet {path} is missing columns {missing_columns}" + ) + for batch in parquet_file.iter_batches(batch_size=1024, columns=columns): + values = batch.to_pydict() + for name, key_id, vector in zip( + values[FEATURE_STORE_PK_FIELD], + values[FEATURE_STORE_SK_FIELD], + values[FEATURE_STORE_VALUE_FIELD], + ): + name = str(name) + if embedding_name is not None and name != embedding_name: + continue + identity = (name, int(key_id)) + if identity in seen: + continue + if vector is None or len(vector) == 0: + raise ValueError( + f"delta parquet {path} contains an empty embedding" + ) + seen.add(identity) + samples.append( + LocalSample( + embedding_name=name, + key_id=int(key_id), + embedding=np.asarray(vector, dtype=np.float32), + source_path=path, + ) + ) + if len(samples) >= sample_count: + return samples + if not samples: + suffix = f" for embedding_name={embedding_name!r}" if embedding_name else "" + raise ValueError(f"no sampleable delta records were found{suffix}") + return samples + + +def _normalize_remote_key(value: Any) -> int: + """Normalize the SDK's string/bytes/integer SK representation.""" + if isinstance(value, bytes): + value = value.decode("utf-8") + return int(value) + + +def verify_samples( + view: Any, + version: str, + samples: Sequence[LocalSample], +) -> Tuple[List[Dict[str, Any]], Dict[str, int]]: + """Query sampled keys and classify presence and value equality.""" + grouped: DefaultDict[str, List[LocalSample]] = defaultdict(list) + for sample in samples: + grouped[sample.embedding_name].append(sample) + + results: List[Dict[str, Any]] = [] + for embedding_name, group in grouped.items(): + remote_rows = view.get_online_features( + feature_name=embedding_name, + keys=[sample.key_id for sample in group], + version=version, + ) + remote_by_key: Dict[int, npt.NDArray[np.float32]] = {} + for row in remote_rows: + raw_key = row.get("sk", row.get(FEATURE_STORE_SK_FIELD)) + if raw_key is None: + raise ValueError("FeatureStore readback row is missing its SK") + key_id = _normalize_remote_key(raw_key) + if key_id in remote_by_key: + raise ValueError( + "FeatureStore returned duplicate rows for " + f"{embedding_name}/{key_id}" + ) + vector = row.get(FEATURE_STORE_VALUE_FIELD) + if vector is None: + raise ValueError( + f"FeatureStore readback row is missing embedding for " + f"{embedding_name}/{key_id}" + ) + remote_by_key[key_id] = np.asarray(vector, dtype=np.float32) + + for sample in group: + remote = remote_by_key.get(sample.key_id) + result: Dict[str, Any] = { + "embedding_name": sample.embedding_name, + "key_id": sample.key_id, + "local_dimension": int(sample.embedding.size), + "source_path": sample.source_path, + } + if remote is None: + result.update( + { + "status": "MISSING", + "remote_dimension": None, + "remote_embedding": None, + } + ) + else: + same_shape = remote.shape == sample.embedding.shape + matches = same_shape and bool( + np.allclose(remote, sample.embedding, rtol=1e-5, atol=1e-6) + ) + max_abs_diff = ( + float(np.max(np.abs(remote - sample.embedding))) + if same_shape and remote.size > 0 + else None + ) + result.update( + { + "status": "MATCH" if matches else "PRESENT_DIFFERENT", + "remote_dimension": int(remote.size), + "remote_embedding": remote.tolist(), + "max_abs_diff": max_abs_diff, + } + ) + results.append(result) + + summary = { + "requested": len(results), + "found": sum(result["status"] != "MISSING" for result in results), + "matching": sum(result["status"] == "MATCH" for result in results), + "present_different": sum( + result["status"] == "PRESENT_DIFFERENT" for result in results + ), + "missing": sum(result["status"] == "MISSING" for result in results), + } + return results, summary + + +def create_feature_store_view(settings: FeatureStoreUploadSettings) -> Any: + """Create the SDK client and return the existing DynamicEmbedding view.""" + try: + from feature_store_py import FeatureStoreClient + except ImportError as exc: + raise RuntimeError( + "feature_store_py is required; install requirements/runtime.txt" + ) from exc + + try: + from alibabacloud_credentials.client import Client as CredClient + except ImportError as exc: + raise RuntimeError( + "alibabacloud_credentials is required; " + "install it via: pip install alibabacloud_credentials" + ) from exc + + credential = CredClient().get_credential() + kwargs = { + "access_key_id": credential.access_key_id, + "access_key_secret": credential.access_key_secret, + "region": settings.region or None, + "endpoint": settings.endpoint or None, + "security_token": credential.security_token or None, + "featuredb_username": os.environ.get("FEATUREDB_USERNAME") or None, + "featuredb_password": os.environ.get("FEATUREDB_PASSWORD") or None, + } + try: + parameters = inspect.signature(FeatureStoreClient).parameters + except (TypeError, ValueError): + parameters = {} + if "test_mode" in parameters: + kwargs["test_mode"] = True + + client = FeatureStoreClient(**kwargs) + project = client.get_project(settings.project_name) + if project is None: + raise RuntimeError( + f"FeatureStore project {settings.project_name!r} was not found" + ) + view = project.get_dynamic_embedding_feature_view(settings.feature_view_name) + if view is None: + raise RuntimeError( + f"DynamicEmbedding FeatureView {settings.feature_view_name!r} was not found" + ) + actual_fields = (view.pk_field, view.sk_field, view.embedding_field) + expected_fields = ( + FEATURE_STORE_PK_FIELD, + FEATURE_STORE_SK_FIELD, + FEATURE_STORE_VALUE_FIELD, + ) + if actual_fields != expected_fields: + raise RuntimeError( + "DynamicEmbedding FeatureView schema mismatch: " + f"expected={expected_fields}, actual={actual_fields}" + ) + return view + + +def run_check(args: argparse.Namespace) -> int: + """Run one local-parquet plus remote-readback verification.""" + pipeline_config = config_util.load_pipeline_config(args.pipeline_config) + train_config = pipeline_config.train_config + if not train_config.HasField("delta_embedding_dump_config"): + raise ValueError("pipeline config has no delta_embedding_dump_config") + dump_config = train_config.delta_embedding_dump_config + if not dump_config.HasField("feature_store_config"): + raise ValueError( + "pipeline config delta_embedding_dump_config has no feature_store_config" + ) + feature_store_config = dump_config.feature_store_config + if not feature_store_config.retain_local_dump: + raise ValueError( + "feature_store_config.retain_local_dump must be enabled so the " + "training job keeps local delta parquet files for readback" + ) + settings = FeatureStoreUploadSettings.from_proto(feature_store_config) + output_dir = resolve_output_dir( + args.pipeline_config, + pipeline_config.model_dir, + dump_config.output_dir, + args.output_dir, + ) + if not os.path.isdir(output_dir): + raise FileNotFoundError( + f"delta embedding output directory not found: {output_dir}" + ) + + file_prefix = dump_config.file_prefix or "delta_embedding" + world_size = args.world_size + global_step, parquet_paths = resolve_upload_step( + output_dir, file_prefix, world_size, args.global_step + ) + samples = sample_local_records( + parquet_paths, + args.sample_count, + embedding_name=args.embedding_name, + ) + + view = create_feature_store_view(settings) + try: + results, summary = verify_samples(view, settings.version, samples) + finally: + close = getattr(view, "close", None) + if callable(close): + close(wait=True) + + report: Dict[str, Any] = { + "target": { + "project_name": settings.project_name, + "feature_view_name": settings.feature_view_name, + "version": settings.version, + }, + "parquet_source": { + "global_step": global_step, + "parquet_paths": [ + os.path.relpath(path, output_dir) for path in parquet_paths + ], + }, + "summary": summary, + "presence_verified": summary["missing"] == 0, + "value_match_verified": summary["matching"] == summary["requested"], + "samples": results, + } + if summary["present_different"]: + report["value_match_note"] = ( + "PRESENT_DIFFERENT confirms the key exists but its value differs from " + "the sampled parquet; a later upload may have updated the same key." + ) + print(json.dumps(report, indent=2, sort_keys=True)) + + if summary["missing"]: + return 1 + if args.require_value_match and summary["present_different"]: + return 1 + return 0 + + +def parse_args(argv: Optional[Sequence[str]] = None) -> argparse.Namespace: + """Parse command-line arguments.""" + parser = argparse.ArgumentParser( + description=( + "Sample a committed delta parquet and read the same keys from " + "FeatureStore using the configured explicit version." + ) + ) + parser.add_argument("--pipeline_config", required=True) + parser.add_argument( + "--output_dir", + default=None, + help="Override delta_embedding_dump output_dir (useful across mounts).", + ) + parser.add_argument( + "--global_step", + type=int, + default=None, + help="Step to inspect; defaults to the latest available shard set.", + ) + parser.add_argument( + "--world_size", + type=int, + default=1, + help="Number of rank shards per step (default: 1 for single-rank).", + ) + parser.add_argument( + "--sample_count", + type=int, + default=10, + help="Maximum number of local keys to read back (default: 10).", + ) + parser.add_argument( + "--embedding_name", + default=None, + help="Only sample rows for this canonical embedding name.", + ) + parser.add_argument( + "--require_value_match", + action="store_true", + help="Exit nonzero when a key exists but differs from the sampled local value.", + ) + return parser.parse_args(argv) + + +def main() -> None: + """CLI entry point.""" + try: + exit_code = run_check(parse_args()) + except Exception as exc: + print(f"ERROR: {exc}", file=sys.stderr) + exit_code = 2 + raise SystemExit(exit_code) + + +if __name__ == "__main__": + main() diff --git a/tzrec/tools/feature_store/check_feature_store_delta_test.py b/tzrec/tools/feature_store/check_feature_store_delta_test.py new file mode 100644 index 00000000..e62b8ac2 --- /dev/null +++ b/tzrec/tools/feature_store/check_feature_store_delta_test.py @@ -0,0 +1,249 @@ +# Copyright (c) 2026, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import sys +import tempfile +import unittest +from types import SimpleNamespace +from unittest import mock + +import numpy as np +import pyarrow as pa +import pyarrow.parquet as pq + +from tzrec.tools.feature_store.check_feature_store_delta import ( + LocalSample, + create_feature_store_view, + parse_args, + resolve_output_dir, + resolve_upload_step, + sample_local_records, + verify_samples, +) + + +class _FakeView: + def get_online_features(self, feature_name, keys, version): + self.feature_name = feature_name + self.keys = keys + self.version = version + return [ + {"sk": str(keys[0]), "embedding": [1.0, 2.0]}, + {"sk": str(keys[1]), "embedding": [9.0, 9.0]}, + ] + + +class CheckFeatureStoreDeltaTest(unittest.TestCase): + def test_create_feature_store_view_uses_public_endpoint_when_supported(self): + captured_kwargs = {} + view = SimpleNamespace( + pk_field="embedding_name", + sk_field="key_id", + embedding_field="embedding", + ) + project = SimpleNamespace(get_dynamic_embedding_feature_view=lambda name: view) + + class FakeFeatureStoreClient: + def __init__(self, test_mode=False, **kwargs): + captured_kwargs.update(kwargs) + captured_kwargs["test_mode"] = test_mode + + def get_project(self, name): + return project + + class FakeCredClient: + def get_credential(self): + return SimpleNamespace( + access_key_id="fake-ak", + access_key_secret="fake-sk", + security_token="fake-sts", + ) + + settings = SimpleNamespace( + region="cn-test", + endpoint="", + project_name="project", + feature_view_name="view", + ) + feature_store_module = SimpleNamespace( + FeatureStoreClient=FakeFeatureStoreClient + ) + cred_module = SimpleNamespace(Client=FakeCredClient) + + with mock.patch.dict( + sys.modules, + { + "feature_store_py": feature_store_module, + "alibabacloud_credentials": SimpleNamespace(client=cred_module), + "alibabacloud_credentials.client": cred_module, + }, + ): + actual = create_feature_store_view(settings) + + self.assertIs(actual, view) + self.assertTrue(captured_kwargs["test_mode"]) + + def test_parse_args_does_not_accept_credentials(self): + args = parse_args(["--pipeline_config", "pipeline.config"]) + + self.assertEqual(args.pipeline_config, "pipeline.config") + self.assertFalse(hasattr(args, "access_key_id")) + self.assertFalse(hasattr(args, "access_key_secret")) + self.assertFalse(hasattr(args, "featuredb_username")) + self.assertFalse(hasattr(args, "featuredb_password")) + + def test_resolve_output_dir_uses_colocated_relocated_outbox(self): + with tempfile.TemporaryDirectory() as tmp_dir: + config_path = os.path.join(tmp_dir, "pipeline.config") + colocated = os.path.join(tmp_dir, "delta_embedding_dump") + os.mkdir(colocated) + self.assertEqual( + resolve_output_dir(config_path, "/missing/model", "", None), + colocated, + ) + + def test_resolve_upload_step_finds_latest_single_rank(self): + with tempfile.TemporaryDirectory() as output_dir: + prefix = "delta__fs_target" + for step in (10, 20, 5): + path = os.path.join(output_dir, f"{prefix}_step_{step}.parquet") + pq.write_table( + pa.table( + { + "embedding_name": ["a"], + "key_id": pa.array([1], type=pa.int64()), + "embedding": pa.array([[1.0]], type=pa.list_(pa.float32())), + } + ), + path, + ) + step, paths = resolve_upload_step(output_dir, prefix, world_size=1) + self.assertEqual(step, 20) + self.assertEqual(len(paths), 1) + self.assertIn("step_20", paths[0]) + + def test_resolve_upload_step_finds_latest_multi_rank(self): + with tempfile.TemporaryDirectory() as output_dir: + prefix = "delta__fs_target" + for step in (10, 20): + step_dir = os.path.join(output_dir, f"step_{step}") + os.makedirs(step_dir) + for rank in range(2): + path = os.path.join( + step_dir, + f"{prefix}_step_{step}_rank_{rank}_of_2.parquet", + ) + pq.write_table( + pa.table( + { + "embedding_name": ["a"], + "key_id": pa.array([rank + 1], type=pa.int64()), + "embedding": pa.array( + [[1.0]], type=pa.list_(pa.float32()) + ), + } + ), + path, + ) + step, paths = resolve_upload_step(output_dir, prefix, world_size=2) + self.assertEqual(step, 20) + self.assertEqual(len(paths), 2) + + def test_resolve_upload_step_explicit_step(self): + with tempfile.TemporaryDirectory() as output_dir: + prefix = "delta__fs_target" + path = os.path.join(output_dir, f"{prefix}_step_15.parquet") + pq.write_table( + pa.table( + { + "embedding_name": ["a"], + "key_id": pa.array([1], type=pa.int64()), + "embedding": pa.array([[1.0]], type=pa.list_(pa.float32())), + } + ), + path, + ) + step, paths = resolve_upload_step( + output_dir, prefix, world_size=1, global_step=15 + ) + self.assertEqual(step, 15) + self.assertEqual(paths, [path]) + + def test_resolve_upload_step_raises_when_not_found(self): + with tempfile.TemporaryDirectory() as output_dir: + with self.assertRaises(FileNotFoundError): + resolve_upload_step(output_dir, "delta__fs_target", world_size=1) + + def test_parquet_paths_and_sampling(self): + with tempfile.TemporaryDirectory() as output_dir: + prefix = "delta__fs_target" + path = os.path.join(output_dir, f"{prefix}_step_20.parquet") + pq.write_table( + pa.table( + { + "embedding_name": ["table_a", "table_a", "table_b"], + "key_id": pa.array([1, 2, 3], type=pa.int64()), + "embedding": pa.array( + [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], + type=pa.list_(pa.float32()), + ), + } + ), + path, + ) + step, paths = resolve_upload_step( + output_dir, prefix, world_size=1, global_step=20 + ) + samples = sample_local_records(paths, 2) + + self.assertEqual(step, 20) + self.assertEqual(paths, [path]) + self.assertEqual( + [(sample.embedding_name, sample.key_id) for sample in samples], + [("table_a", 1), ("table_a", 2)], + ) + np.testing.assert_array_equal(samples[0].embedding, [1.0, 2.0]) + + def test_verify_samples_separates_presence_from_value_match(self): + samples = [ + LocalSample("table_a", 1, np.array([1.0, 2.0], np.float32), "a"), + LocalSample("table_a", 2, np.array([3.0, 4.0], np.float32), "a"), + LocalSample("table_a", 3, np.array([5.0, 6.0], np.float32), "a"), + ] + view = _FakeView() + + results, summary = verify_samples(view, "v1", samples) + + self.assertEqual(view.feature_name, "table_a") + self.assertEqual(view.keys, [1, 2, 3]) + self.assertEqual(view.version, "v1") + self.assertEqual( + [result["status"] for result in results], + ["MATCH", "PRESENT_DIFFERENT", "MISSING"], + ) + self.assertEqual(results[0]["remote_embedding"], [1.0, 2.0]) + self.assertEqual(results[1]["remote_embedding"], [9.0, 9.0]) + self.assertIsNone(results[2]["remote_embedding"]) + self.assertEqual( + summary, + { + "requested": 3, + "found": 2, + "matching": 1, + "present_different": 1, + "missing": 1, + }, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/utils/delta_embedding_dump.py b/tzrec/utils/delta_embedding_dump.py index a00ddd46..fbea3ee3 100644 --- a/tzrec/utils/delta_embedding_dump.py +++ b/tzrec/utils/delta_embedding_dump.py @@ -10,6 +10,7 @@ # limitations under the License. import os +import time from contextlib import contextmanager from dataclasses import dataclass from typing import ( @@ -53,6 +54,7 @@ from tzrec.protos.feature_pb2 import FeatureConfig from tzrec.protos.train_pb2 import DeltaEmbeddingDumpConfig +from tzrec.utils.feature_store_delta_uploader import FeatureStoreDeltaUploader from tzrec.utils.logging_util import logger _CONSUMER = "delta_embedding_dump" @@ -116,7 +118,17 @@ def validate_delta_embedding_dump_config( "delta_embedding_dump_config only supports CUDA training, " f"but got device={device}." ) - if config.dump_interval_steps <= 0: + if config.HasField("dump_interval_minutes"): + if config.HasField("dump_interval_steps"): + raise ValueError( + "delta_embedding_dump_config must configure only one of " + "dump_interval_steps and dump_interval_minutes." + ) + if config.dump_interval_minutes <= 0: + raise ValueError( + "delta_embedding_dump_config.dump_interval_minutes must be > 0." + ) + elif config.dump_interval_steps <= 0: raise ValueError("delta_embedding_dump_config.dump_interval_steps must be > 0.") @@ -565,13 +577,24 @@ def __init__( validate_delta_embedding_dump_no_zch_features(feature_configs) self._model = model self._config = config - self._interval = config.dump_interval_steps + self._interval_steps: Optional[int] = None + self._interval_secs: Optional[float] = None + if config.HasField("dump_interval_minutes"): + self._interval_secs = float(config.dump_interval_minutes * 60) + else: + self._interval_steps = int(config.dump_interval_steps) + self._next_dump_time: Optional[float] = None + self._last_dump_step: Optional[int] = None self._output_dir = config.output_dir or os.path.join( model_dir, "delta_embedding_dump" ) self._file_prefix = config.file_prefix or "delta_embedding" self._rank, self._world_size = _distributed_rank_world_size() self._tracking_pause_depth = 0 + self._feature_store_enabled = config.HasField("feature_store_config") + self._retain_local_dump = self._feature_store_enabled and bool( + config.feature_store_config.retain_local_dump + ) os.makedirs(self._output_dir, exist_ok=True) self._tracker = ModelDeltaTracker( @@ -583,14 +606,35 @@ def __init__( self._table_shard_infos = self._collect_table_shard_infos() self._validate_supported_table_sharding(self._table_shard_infos) self._install_tracking_pause_guard() + self._uploader: Optional[FeatureStoreDeltaUploader] = None + if self._feature_store_enabled: + embedding_dimensions = { + fqn: int(info.global_cols) + for fqn, info in self._table_shard_infos.items() + } + self._uploader = FeatureStoreDeltaUploader( + config.feature_store_config, + embedding_dimensions=embedding_dimensions, + rank=self._rank, + world_size=self._world_size, + manage_remote_view=self._rank == 0, + ) + interval_name = "minutes" if self._interval_secs is not None else "steps" + interval_value = ( + config.dump_interval_minutes + if self._interval_secs is not None + else self._interval_steps + ) logger.info( - "Delta embedding dump enabled: interval=%s output_dir=%s " - "rank=%s/%s tables=%s", - self._interval, + "Delta embedding dump enabled: interval_%s=%s output_dir=%s " + "rank=%s/%s tables=%s feature_store_upload=%s", + interval_name, + interval_value, self._output_dir, self._rank, self._world_size, sorted(self._tracker.fqn_to_feature_names), + self._feature_store_enabled, ) def clear(self) -> None: @@ -606,16 +650,69 @@ def pause_tracking(self) -> Iterator[None]: finally: self._tracking_pause_depth -= 1 + def start(self) -> None: + """Start timed cadence and per-rank FeatureStore publication. + + The rank-zero view rendezvous (create the DynamicEmbedding view, + barrier, then non-primary ranks open it) lives in + :meth:`FeatureStoreDeltaUploader.start`; this method only delegates and + arms the timed cadence. + """ + if self._uploader is not None: + self._uploader.start() + if self._interval_secs is not None: + self._next_dump_time = time.monotonic() + self._interval_secs + + def close(self, raise_on_error: bool = True, drain: bool = True) -> None: + """Close this rank's uploader; abnormal shutdown can skip draining.""" + if self._uploader is not None: + self._uploader.close(raise_on_error=raise_on_error, drain=drain) + + def _feature_store_upload_error(self) -> Optional[BaseException]: + """Collect this rank's uploader error without changing control flow.""" + if not self._feature_store_enabled: + return None + if self._uploader is not None: + try: + self._uploader.check_error() + except BaseException as exc: + return exc + return None + + def _check_feature_store_upload_error(self) -> None: + """Surface this rank's background upload failure to the trainer.""" + error = self._feature_store_upload_error() + if error is not None: + raise error.with_traceback(error.__traceback__) + def maybe_dump(self, global_step: int) -> None: - """Dump on the configured global-step interval and advance tracker state. + """Dump on the configured step or time interval and advance tracker state. Args: global_step: Current training step. """ - if global_step > 0 and global_step % self._interval == 0: + self._check_feature_store_upload_error() + if self._local_dump_decision(global_step): self.dump(global_step) + self._last_dump_step = global_step + if self._interval_secs is not None and self._next_dump_time is not None: + # Fixed-rate rescheduling keeps every rank's deadline sequence + # identical, so timed dumps stay step-aligned across ranks up to + # clock skew exactly astride a step boundary; missed deadlines + # are skipped instead of fired as a burst. + now = time.monotonic() + while self._next_dump_time <= now: + self._next_dump_time += self._interval_secs self._tracker.step() + def _local_dump_decision(self, global_step: int) -> bool: + """Return whether this step triggers a delta dump.""" + if self._interval_steps is not None: + return global_step > 0 and global_step % self._interval_steps == 0 + if self._interval_secs is not None and self._next_dump_time is not None: + return time.monotonic() >= self._next_dump_time + return False + def final_dump(self, global_step: int) -> Optional[str]: """Flush the trailing partial interval at the end of training. @@ -629,15 +726,24 @@ def final_dump(self, global_step: int) -> Optional[str]: Returns: Path to the dumped parquet file, or None if skipped. """ + if global_step <= 0: + logger.info("Skipping delta embedding dump at step %s.", global_step) + return None global_step = self._sync_final_step(global_step) - if global_step > 0 and global_step % self._interval == 0: + if self._interval_steps is not None and global_step % self._interval_steps == 0: # Boundary steps were already written (with full delta) by # ``maybe_dump``. Re-dumping here has no new delta to flush -- every # rank's consumer cursor has already advanced past the boundary's - # delta -- and torchrec's ``get_unique`` raises - # ``torch.cat(): expected a non-empty list of Tensors`` on the empty - # consumer window. Re-dumping would also overwrite the already-written - # boundary shards (with an empty file under multi-GPU), so skip. + # delta (all ranks run the same step count, so every rank + # participated in the boundary dump) -- and torchrec's + # ``get_unique`` raises ``torch.cat(): expected a non-empty list of + # Tensors`` on the empty consumer window. Re-dumping would also + # overwrite the already-written boundary shards (with an empty file + # under multi-GPU), so skip. + return None + if self._interval_secs is not None and global_step == self._last_dump_step: + # A timed dump can land on any step. Avoid replacing that step's full + # delta with an empty final shard when training ends immediately after. return None return self.dump(global_step) @@ -673,6 +779,11 @@ def dump(self, global_step: int) -> Optional[str]: Returns: Path to the dumped parquet file, or None if no data to dump. """ + global_step = int(global_step) + if global_step <= 0: + raise ValueError("delta embedding dump global_step must be > 0") + uploader = self._uploader + write_local = not self._feature_store_enabled or self._retain_local_dump table_weights = self._collect_table_weights() dynamic_modules = self._collect_dynamic_modules() table_chunks: List[pa.Table] = [] @@ -682,21 +793,31 @@ def dump(self, global_step: int) -> Optional[str]: table_weights=table_weights, dynamic_modules=dynamic_modules, ) - if num_rows == 0: - if self._world_size == 1: - logger.info("No delta embedding rows to dump at step %s.", global_step) - return None + output_path: Optional[str] = None + if write_local and (num_rows > 0 or self._world_size > 1): + # Multi-rank shard sets stay complete even for an empty rank so + # per-step file consumers never observe a partial set. output_path = self._output_path(global_step) self._write_table_chunks(table_chunks, output_path) - logger.info( - "Dumped empty delta embedding shard to %s at step %s.", - output_path, + if uploader is not None and num_rows > 0: + uploader.submit(global_step, pa.concat_tables(table_chunks)) + if num_rows == 0: + if output_path is None: + logger.debug("No delta embedding rows to dump at step %s.", global_step) + else: + logger.debug( + "Dumped empty delta embedding shard to %s at step %s.", + output_path, + global_step, + ) + elif output_path is None: + logger.debug( + "Submitted %s delta embedding rows for FeatureStore upload at step %s.", + num_rows, global_step, ) - return output_path - output_path = self._output_path(global_step) - self._write_table_chunks(table_chunks, output_path) - logger.info("Dumped %s delta embedding rows to %s.", num_rows, output_path) + else: + logger.debug("Dumped %s delta embedding rows to %s.", num_rows, output_path) return output_path def _output_path(self, global_step: int) -> str: diff --git a/tzrec/utils/delta_embedding_dump_test.py b/tzrec/utils/delta_embedding_dump_test.py index cbf89890..4272b853 100644 --- a/tzrec/utils/delta_embedding_dump_test.py +++ b/tzrec/utils/delta_embedding_dump_test.py @@ -393,6 +393,15 @@ def test_present_config_requires_positive_interval(self): with self.assertRaisesRegex(ValueError, "dump_interval_steps"): validate_delta_embedding_dump_config(config, torch.device("cuda:0")) + def test_present_config_accepts_minutes_interval(self): + config = DeltaEmbeddingDumpConfig(dump_interval_minutes=5) + validate_delta_embedding_dump_config(config, torch.device("cuda:0")) + + def test_present_config_requires_positive_minutes_interval(self): + config = DeltaEmbeddingDumpConfig(dump_interval_minutes=0) + with self.assertRaisesRegex(ValueError, "dump_interval_minutes"): + validate_delta_embedding_dump_config(config, torch.device("cuda:0")) + def test_zch_feature_fails_fast(self): feature_configs = [ feature_pb2.FeatureConfig( @@ -567,7 +576,9 @@ def test_write_table_chunks_leaves_no_partial_shard_on_error(self): def test_final_dump_skips_boundary_step_to_avoid_overwrite(self): dumper = object.__new__(DeltaEmbeddingDumper) - dumper._interval = 50 + dumper._interval_steps = 50 + dumper._interval_secs = None + dumper._last_dump_step = None dumper._world_size = 1 with mock.patch.object(dumper, "dump") as dump_mock: # Boundary steps were already written by maybe_dump; skip them so a @@ -576,12 +587,13 @@ def test_final_dump_skips_boundary_step_to_avoid_overwrite(self): self.assertIsNone(dumper.final_dump(100)) dump_mock.assert_not_called() - # Trailing partial interval (and step 0) must still be flushed. - dumper.final_dump(0) + # Step 0 is not publishable; final_dump returns early. A positive + # trailing partial interval must still be flushed. + self.assertIsNone(dumper.final_dump(0)) dumper.final_dump(73) self.assertEqual( [call.args[0] for call in dump_mock.call_args_list], - [0, 73], + [73], ) def test_final_dump_syncs_step_across_ranks_before_flush(self): @@ -590,7 +602,9 @@ def test_final_dump_syncs_step_across_ranks_before_flush(self): # skip and write no shard, leaving step_73/ ragged. The MAX all_reduce # lifts every rank to 73 so all take the same dump-into-step_73 path. dumper = object.__new__(DeltaEmbeddingDumper) - dumper._interval = 50 + dumper._interval_steps = 50 + dumper._interval_secs = None + dumper._last_dump_step = None dumper._world_size = 2 def fake_all_reduce(tensor, op=None): @@ -611,9 +625,25 @@ def fake_all_reduce(tensor, op=None): dumper.final_dump(50) dump_mock.assert_called_once_with(73) + def test_final_dump_skips_step_already_dumped_by_time_interval(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._interval_steps = None + dumper._interval_secs = 60.0 + dumper._last_dump_step = 73 + dumper._world_size = 1 + with mock.patch.object(dumper, "dump") as dump_mock: + self.assertIsNone(dumper.final_dump(73)) + dump_mock.assert_not_called() + def test_maybe_dump_uses_checkpoint_aligned_global_step(self): dumper = object.__new__(DeltaEmbeddingDumper) - dumper._interval = 50 + dumper._interval_steps = 50 + dumper._interval_secs = None + dumper._last_dump_step = None + dumper._rank = 0 + dumper._world_size = 1 + dumper._feature_store_enabled = False + dumper._uploader = None dumper._tracker = mock.MagicMock() with mock.patch.object(dumper, "dump") as dump_mock: dumper.maybe_dump(49) @@ -629,6 +659,138 @@ def test_maybe_dump_uses_checkpoint_aligned_global_step(self): ) self.assertEqual(dumper._tracker.step.call_count, 4) + def test_maybe_dump_uses_elapsed_time_with_fixed_rate_schedule(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._interval_steps = None + dumper._interval_secs = 60.0 + dumper._next_dump_time = 160.0 + dumper._last_dump_step = None + dumper._rank = 0 + dumper._world_size = 1 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._tracker = mock.MagicMock() + with ( + mock.patch.object(dumper, "dump") as dump_mock, + mock.patch( + "tzrec.utils.delta_embedding_dump.time.monotonic", + side_effect=[159.0, 160.0, 162.0, 221.0, 222.0, 223.0], + ), + ): + dumper.maybe_dump(10) + dumper.maybe_dump(11) + dumper.maybe_dump(12) + dumper.maybe_dump(13) + + self.assertEqual( + [call.args[0] for call in dump_mock.call_args_list], + [11, 12], + ) + # Deadlines advance at a fixed rate from the armed schedule + # (160 -> 220 -> 280), not from each dump's completion time. + self.assertEqual(dumper._next_dump_time, 280.0) + self.assertEqual(dumper._last_dump_step, 12) + self.assertEqual(dumper._tracker.step.call_count, 4) + + def test_timed_dump_decides_locally_without_collectives(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._interval_steps = None + dumper._interval_secs = 60.0 + dumper._next_dump_time = 160.0 + dumper._last_dump_step = None + dumper._rank = 1 + dumper._world_size = 2 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._tracker = mock.MagicMock() + with ( + mock.patch.object( + dumper, "dump", return_value="delta.parquet" + ) as dump_mock, + mock.patch("torch.distributed.all_reduce") as all_reduce_mock, + mock.patch( + "tzrec.utils.delta_embedding_dump.time.monotonic", + side_effect=[159.0, 160.5, 161.0], + ), + ): + dumper.maybe_dump(10) + dump_mock.assert_not_called() + dumper.maybe_dump(11) + + all_reduce_mock.assert_not_called() + dump_mock.assert_called_once_with(11) + self.assertEqual(dumper._last_dump_step, 11) + self.assertEqual(dumper._next_dump_time, 220.0) + self.assertEqual(dumper._tracker.step.call_count, 2) + + def test_timed_maybe_dump_propagates_local_dump_failure(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._interval_steps = None + dumper._interval_secs = 60.0 + dumper._next_dump_time = 0.0 + dumper._last_dump_step = None + dumper._rank = 0 + dumper._world_size = 2 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._tracker = mock.MagicMock() + dump_error = RuntimeError("local dump failed") + with ( + mock.patch.object(dumper, "dump", side_effect=dump_error) as dump_mock, + mock.patch( + "tzrec.utils.delta_embedding_dump.time.monotonic", return_value=1.0 + ), + ): + with self.assertRaises(RuntimeError) as context: + dumper.maybe_dump(10) + + self.assertIs(context.exception, dump_error) + dump_mock.assert_called_once_with(10) + self.assertIsNone(dumper._last_dump_step) + self.assertEqual(dumper._tracker.step.call_count, 0) + + def test_timed_dump_skips_missed_deadlines_without_burst(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._interval_steps = None + dumper._interval_secs = 60.0 + dumper._next_dump_time = 0.0 + dumper._last_dump_step = None + dumper._rank = 0 + dumper._world_size = 2 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._tracker = mock.MagicMock() + with ( + mock.patch.object( + dumper, "dump", return_value="delta.parquet" + ) as dump_mock, + mock.patch( + "tzrec.utils.delta_embedding_dump.time.monotonic", + return_value=100.0, + ), + ): + dumper.maybe_dump(10) + dumper.maybe_dump(11) + dumper.maybe_dump(12) + + dump_mock.assert_called_once_with(10) + # Deadlines 0 and 60 already elapsed at the dump; skip past them + # instead of firing a burst of catch-up dumps. + self.assertEqual(dumper._next_dump_time, 120.0) + self.assertEqual(dumper._tracker.step.call_count, 3) + + def test_start_initializes_minutes_interval_from_training_start(self): + dumper = object.__new__(DeltaEmbeddingDumper) + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._interval_secs = 120.0 + dumper._next_dump_time = None + with mock.patch( + "tzrec.utils.delta_embedding_dump.time.monotonic", return_value=100.0 + ): + dumper.start() + self.assertEqual(dumper._next_dump_time, 220.0) + def test_tracker_uses_auto_compact(self): tracker = mock.MagicMock() tracker.fqn_to_feature_names = {} @@ -650,6 +812,28 @@ def test_tracker_uses_auto_compact(self): self.assertTrue(tracker_cls.call_args.kwargs["auto_compact"]) + def test_minutes_interval_is_converted_to_seconds(self): + tracker = mock.MagicMock() + tracker.fqn_to_feature_names = {} + tracker.tracked_modules = {} + with ( + tempfile.TemporaryDirectory() as tmp_dir, + mock.patch( + "tzrec.utils.delta_embedding_dump.ModelDeltaTracker", + return_value=tracker, + ), + ): + dumper = DeltaEmbeddingDumper( + torch.nn.Module(), + DeltaEmbeddingDumpConfig(dump_interval_minutes=2), + tmp_dir, + torch.device("cuda"), + [], + ) + + self.assertIsNone(dumper._interval_steps) + self.assertEqual(dumper._interval_secs, 120.0) + def test_model_delta_tracker_records_same_table_name_by_owner_fqn(self): tracker = object.__new__(ModelDeltaTracker) ebc_module = torch.nn.Module() @@ -918,6 +1102,9 @@ def test_multi_gpu_dump_writes_empty_shard_when_rank_has_no_delta(self): dumper._file_prefix = "delta_embedding" dumper._rank = 1 dumper._world_size = 2 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._retain_local_dump = False with ( mock.patch.object(dumper, "_collect_table_weights", return_value={}), mock.patch.object(dumper, "_collect_dynamic_modules", return_value={}), @@ -944,6 +1131,9 @@ def test_single_gpu_dump_skips_file_when_rank_has_no_delta(self): dumper._file_prefix = "delta_embedding" dumper._rank = 0 dumper._world_size = 1 + dumper._feature_store_enabled = False + dumper._uploader = None + dumper._retain_local_dump = False with ( mock.patch.object(dumper, "_collect_table_weights", return_value={}), mock.patch.object(dumper, "_collect_dynamic_modules", return_value={}), diff --git a/tzrec/utils/dist_util.py b/tzrec/utils/dist_util.py index bc916a7f..88bb0f18 100644 --- a/tzrec/utils/dist_util.py +++ b/tzrec/utils/dist_util.py @@ -288,7 +288,8 @@ def _next_batch(self, dataloader_iter: Iterator[In]) -> Optional[In]: 0 if batch is None else 1, dtype=torch.float, device=self._device ) dist.all_reduce(has_batch, dist.ReduceOp.AVG) - if has_batch.item() < 1: + available_fraction = has_batch.item() + if available_fraction < 1: # We drop remainder batches on all workers, # if one worker does not have a batch self._dataloader_exhausted = True diff --git a/tzrec/utils/export_util_test.py b/tzrec/utils/export_util_test.py index 1279f829..570bd690 100644 --- a/tzrec/utils/export_util_test.py +++ b/tzrec/utils/export_util_test.py @@ -40,7 +40,7 @@ from tzrec.protos import feature_pb2, loss_pb2, model_pb2, module_pb2 from tzrec.protos.models import rank_model_pb2 from tzrec.protos.pipeline_pb2 import EasyRecConfig -from tzrec.utils import checkpoint_util, misc_util +from tzrec.utils import checkpoint_util, config_util, misc_util from tzrec.utils.export_util import ( _dedup_key_files_by_realpath, _get_dense_embedding_leaf_module_names, @@ -312,6 +312,139 @@ def forward(self, data, device=None): # type: ignore[no-untyped-def] _restore_env(old_env) shutil.rmtree(tmp, ignore_errors=True) + def test_distributed_embedding_export_uses_overrides_and_preserves_config( + self, + ) -> None: + class FakeBatch: + def to(self, device): # type: ignore[no-untyped-def] + return self + + def to_dict(self, sparse_dtype): # type: ignore[no-untyped-def] + return {"x": torch.ones(1)} + + class FakeDataloader: + dataset = SimpleNamespace(sampled_batch_size=1) + + def __iter__(self): # type: ignore[no-untyped-def] + return iter([FakeBatch()]) + + class TinyModel(torch.nn.Module): + def __init__(self): # type: ignore[no-untyped-def] + super().__init__() + self.features = [] + + def set_is_inference(self, is_inference): # type: ignore[no-untyped-def] + self.is_inference = is_inference + + def forward(self, data, device=None): # type: ignore[no-untyped-def] + return {"score": data["x"] + 1} + + class FakeDMP(torch.nn.Module): + def __init__(self, module, *args, **kwargs): # type: ignore[no-untyped-def] + super().__init__() + self.module = module + + def forward(self, data, device=None): # type: ignore[no-untyped-def] + return self.module(data, device=device) + + tmp = tempfile.mkdtemp(prefix="tzrec_export_dist_overrides_") + old_env = { + key: os.environ.get(key) + for key in ("RANK", "LOCAL_RANK", "WORLD_SIZE", "LOCAL_WORLD_SIZE") + } + try: + os.environ["RANK"] = "0" + os.environ["LOCAL_RANK"] = "0" + os.environ["WORLD_SIZE"] = "1" + os.environ["LOCAL_WORLD_SIZE"] = "1" + pipeline_config = EasyRecConfig( + train_input_path="train_input", + eval_input_path="eval_input", + model_dir="model_dir", + ) + dump_config = pipeline_config.train_config.delta_embedding_dump_config + feature_store_config = dump_config.feature_store_config + feature_store_config.region = "cn-test" + feature_store_config.project_name = "project_a" + feature_store_config.feature_view_name = "shared_embeddings" + feature_store_config.version = "model_a@export_1" + model_acc = {"SPARSE_INT64": "1", "cand_seq_pk": "cand_seq"} + fake_scripted = mock.Mock() + + with ( + mock.patch( + "tzrec.utils.export_util.init_process_group", + return_value=(torch.device("cpu"), None), + ), + mock.patch( + "tzrec.utils.export_util._get_sparse_table_to_embedding_info", + return_value=({}, {}), + ), + mock.patch( + "tzrec.utils.export_util.create_dataloader", + return_value=FakeDataloader(), + ) as create_dataloader_mock, + mock.patch( + "tzrec.utils.export_util.create_planner", + return_value=SimpleNamespace(collective_plan=lambda *args: None), + ), + mock.patch( + "tzrec.utils.export_util.get_default_sharders", return_value=[] + ), + mock.patch( + "tzrec.utils.export_util.DistributedModelParallel", + side_effect=lambda *args, **kwargs: FakeDMP(kwargs["module"]), + ), + mock.patch("tzrec.utils.export_util.checkpoint_util.restore_model"), + mock.patch("tzrec.utils.export_util.init_parameters"), + mock.patch( + "tzrec.utils.export_util._get_sparse_embedding_tensor", + return_value=({}, {}, {}, {}), + ), + mock.patch( + "tzrec.utils.export_util.create_fg_json", + return_value={"features": []}, + ), + mock.patch( + "tzrec.utils.export_util.symbolic_trace", + return_value=SimpleNamespace(code="def forward(self):\n pass\n"), + ), + mock.patch( + "tzrec.utils.export_util.torch.jit.script", + return_value=fake_scripted, + ), + mock.patch( + "tzrec.utils.export_util.acc_utils.export_acc_config", + return_value=model_acc, + ) as export_acc_config_mock, + ): + export_distributed_embedding( + pipeline_config, + TinyModel(), + "checkpoint_dir", + tmp, + additional_export_config={"cand_seq_pk": "cand_seq"}, + data_input_path="override_input", + ) + + create_dataloader_mock.assert_called_once() + self.assertEqual(create_dataloader_mock.call_args.args[2], "override_input") + export_acc_config_mock.assert_called_once_with( + additional_export_config={"cand_seq_pk": "cand_seq"} + ) + with open(os.path.join(tmp, "model_acc.json")) as f: + self.assertEqual(json.load(f), model_acc) + pipeline_config_path = os.path.join(tmp, "pipeline.config") + exported_config = config_util.load_pipeline_config(pipeline_config_path) + exported_dump_config = ( + exported_config.train_config.delta_embedding_dump_config + ) + exported_feature_store_config = exported_dump_config.feature_store_config + self.assertEqual(exported_feature_store_config.project_name, "project_a") + finally: + _restore_env(old_env) + shutil.rmtree(tmp, ignore_errors=True) + def test_sparse_dynamic_embedding_export_concats_training_shards(self) -> None: """Single-rank export must not drop multi-GPU dynamicemb checkpoint shards.""" tmp = tempfile.mkdtemp(prefix="tzrec_export_dynemb_") diff --git a/tzrec/utils/feature_store_delta_uploader.py b/tzrec/utils/feature_store_delta_uploader.py new file mode 100644 index 00000000..7ab0019f --- /dev/null +++ b/tzrec/utils/feature_store_delta_uploader.py @@ -0,0 +1,910 @@ +# Copyright (c) 2026, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Ephemeral per-rank uploader for in-memory delta-embedding tables. + +Best-effort upload for the current live training process only. No cross-restart +recovery, no durable state, no replay. A process crash means restart from the +latest checkpoint and pending deltas are discarded. +""" + +import os +import threading +import time +from collections import deque +from dataclasses import dataclass +from typing import ( + Any, + Callable, + Deque, + Dict, + List, + Mapping, + Optional, + Tuple, + cast, +) + +import numpy as np +import pyarrow as pa + +from tzrec.protos.train_pb2 import FeatureStoreConfig +from tzrec.utils.checkpoint_util import remap_input_tile_user_key +from tzrec.utils.logging_util import logger +from tzrec.utils.sparse_embedding_contract import SPARSE_EMBEDDING_INVALID_KEY + +FEATURE_STORE_PK_FIELD = "embedding_name" +FEATURE_STORE_SK_FIELD = "key_id" +FEATURE_STORE_VALUE_FIELD = "embedding" +FEATURE_STORE_WRITE_MODE = "MERGE" +FEATURE_STORE_SDK_BATCH_SIZE = 1000 +FEATURE_STORE_UPLOAD_FORMAT_DEFAULT = "ARROW" +FEATURE_STORE_UPLOAD_FORMATS = ("ARROW", "JSON") + +# Default FeatureStore entity used when a DynamicEmbedding FeatureView has to be +# created. The entity is provisioned on demand by the rank-zero uploader, so +# users never configure it; the join_id only needs to be a stable non-empty key. +FEATURE_STORE_DEFAULT_ENTITY_NAME = "default_dynemb_entity" +FEATURE_STORE_DEFAULT_ENTITY_JOIN_ID = "default_dynemb_join_id" + +_FEATURE_STORE_PROGRESS_LOG_INTERVAL_BATCHES = 100 + + +@dataclass(frozen=True) +class _DeltaBatch: + """Decoded view of one delta RecordBatch shared by both upload paths. + + PK values are pre-remapped (INPUT_TILE=3 user-side keys -> non-user twin + keys); the raw table_fqn is retained only for contract/dimension checks. + """ + + num_rows: int + remapped_fqns: List[str] + key_ids: np.ndarray + embedding_column: pa.Array + flat_embeddings: np.ndarray + offsets: np.ndarray + + +class FeatureStoreUploadError(RuntimeError): + """Safe, credential-free error propagated from the uploader thread.""" + + +class _UploadAborted(RuntimeError): + """Internal control flow for abnormal, non-draining shutdown.""" + + +@dataclass(frozen=True) +class FeatureStoreUploadSettings: + """Validated immutable settings copied from the runtime protobuf.""" + + region: str + endpoint: str + project_name: str + feature_view_name: str + feature_view_ttl_secs: int + feature_view_shard_count: int + feature_view_replication_count: int + version: str + upload_batch_size: int + max_retries: int + retry_backoff_secs: int + shutdown_timeout_secs: int + max_pending_steps: int + poll_interval_secs: int + upload_format: str + + @classmethod + def from_proto(cls, config: FeatureStoreConfig) -> "FeatureStoreUploadSettings": + """Validate configuration without resolving credentials.""" + initialization_errors = config.FindInitializationErrors() + if initialization_errors: + raise ValueError( + "feature_store_config is missing required fields: " + + ", ".join(initialization_errors) + ) + region = config.region or os.environ.get("ALIBABA_CLOUD_REGION", "") + endpoint = config.endpoint + + if not region: + raise ValueError( + "feature_store_config.region must not be empty " + "(it may come from ALIBABA_CLOUD_REGION)" + ) + project_name = config.project_name.strip() + feature_view_name = config.feature_view_name.strip() + version = config.version.strip() + if not project_name: + raise ValueError("feature_store_config.project_name must not be empty") + if not feature_view_name: + raise ValueError("feature_store_config.feature_view_name must not be empty") + if not version or version == "default": + raise ValueError( + "feature_store_config.version must be an explicit non-default version" + ) + + positive_values = { + "feature_view_ttl_secs": int(config.feature_view_ttl_secs), + "upload_batch_size": int(config.upload_batch_size), + "max_retries": int(config.max_retries), + "shutdown_timeout_secs": int(config.shutdown_timeout_secs), + "max_pending_steps": int(config.max_pending_steps), + "poll_interval_secs": int(config.poll_interval_secs), + } + for name, value in positive_values.items(): + if value <= 0: + raise ValueError(f"feature_store_config.{name} must be > 0") + feature_view_shard_count = int(config.feature_view_shard_count) + if not 1 <= feature_view_shard_count <= 20: + raise ValueError( + "feature_store_config.feature_view_shard_count must be in [1, 20]" + ) + feature_view_replication_count = int(config.feature_view_replication_count) + if not 1 <= feature_view_replication_count <= 3: + raise ValueError( + "feature_store_config.feature_view_replication_count must be in [1, 3]" + ) + if positive_values["upload_batch_size"] > FEATURE_STORE_SDK_BATCH_SIZE: + raise ValueError( + "feature_store_config.upload_batch_size must be <= " + f"{FEATURE_STORE_SDK_BATCH_SIZE} so one publish timestamp maps to " + "exactly one FeatureStore SDK HTTP batch" + ) + + upload_format = ( + (config.upload_format or FEATURE_STORE_UPLOAD_FORMAT_DEFAULT) + .strip() + .upper() + ) + if upload_format not in FEATURE_STORE_UPLOAD_FORMATS: + raise ValueError( + "feature_store_config.upload_format must be one of " + f"{FEATURE_STORE_UPLOAD_FORMATS}, got {upload_format!r}" + ) + return cls( + region=region, + endpoint=endpoint, + project_name=project_name, + feature_view_name=feature_view_name, + feature_view_ttl_secs=positive_values["feature_view_ttl_secs"], + feature_view_shard_count=feature_view_shard_count, + feature_view_replication_count=feature_view_replication_count, + version=version, + upload_batch_size=positive_values["upload_batch_size"], + max_retries=positive_values["max_retries"], + retry_backoff_secs=int(config.retry_backoff_secs), + shutdown_timeout_secs=positive_values["shutdown_timeout_secs"], + max_pending_steps=positive_values["max_pending_steps"], + poll_interval_secs=positive_values["poll_interval_secs"], + upload_format=upload_format, + ) + + +class FeatureStoreDeltaUploader: + """Ephemeral per-rank uploader for in-memory delta-embedding tables. + + Best-effort upload for the current live process only. No cross-restart + recovery, no durable state, no replay. A process crash means restart from + the latest checkpoint and pending deltas are discarded. + + Every rank owns one uploader and streams only its local shard rows, so no + cross-rank aggregation or deduplication is needed: table-wise and row-wise + sharding give each (embedding_name, key_id) a unique owner rank, while + data-parallel replicas issue identical idempotent MERGE writes. Publish + timestamps are allocated per rank and only need per-rank monotonicity, + because a key's owner rank is fixed for the lifetime of the process. + """ + + def __init__( + self, + config: FeatureStoreConfig, + embedding_dimensions: Mapping[str, int], + rank: int = 0, + world_size: int = 1, + manage_remote_view: bool = True, + clock_ms: Optional[Callable[[], int]] = None, + ) -> None: + """Initialize the uploader with validated settings and in-memory state.""" + self._settings = FeatureStoreUploadSettings.from_proto(config) + self._rank = int(rank) + self._world_size = int(world_size) + self._manage_remote_view = bool(manage_remote_view) + self._embedding_dimensions = { + str(name): int(dimension) + for name, dimension in embedding_dimensions.items() + } + + self._credentials_client = self._create_credentials_client() + self._clock_ms = clock_ms or (lambda: time.time_ns() // 1_000_000) + self._view = None + self._condition = threading.Condition() + self._pending: Deque[Tuple[int, pa.Table]] = deque() + self._started = False + self._closing = False + self._aborting = False + self._closed = False + self._worker: Optional[threading.Thread] = None + self._error: Optional[FeatureStoreUploadError] = None + self._last_publish_ts: int = 0 + + def start(self) -> None: + """Start the worker after the rank-zero view rendezvous. + + Rank zero creates and validates the DynamicEmbedding view first; every + rank then barriers so non-primary ranks open the already-published view + without control-plane races. A rank-zero startup failure still joins the + barrier before re-raising so every rank issues the same collective order. + """ + with self._condition: + if self._started: + return + if self._closed: + raise RuntimeError("FeatureStoreDeltaUploader is already closed") + self._raise_if_failed_locked() + start_error: Optional[BaseException] = None + if self._manage_remote_view: + try: + self._get_view() + except BaseException as exc: + self._reset_view(suppress_errors=True) + start_error = exc + if self._world_size > 1: + import torch.distributed as dist + + if not (dist.is_available() and dist.is_initialized()): + raise RuntimeError( + "distributed FeatureStore delta dump requires an " + "initialized process group" + ) + dist.barrier() + if start_error is not None: + raise start_error.with_traceback(start_error.__traceback__) + try: + self._get_view() + except BaseException: + self._reset_view(suppress_errors=True) + raise + self._started = True + self._worker = threading.Thread( + target=self._run, + name="tzrec-feature-store-delta-uploader", + daemon=True, + ) + self._worker.start() + logger.info( + "FeatureStore delta uploader started: project=%s feature_view=%s " + "version=%s rank=%s", + self._settings.project_name, + self._settings.feature_view_name, + self._settings.version, + self._rank, + ) + + def submit(self, global_step: int, table: pa.Table) -> None: + """Enqueue one step's in-memory delta table with back-pressure.""" + global_step = int(global_step) + if global_step <= 0: + raise ValueError("FeatureStore delta global_step must be > 0") + with self._condition: + self._raise_if_failed_locked() + if not self._started: + raise RuntimeError( + "FeatureStoreDeltaUploader.start() must be called before submit()" + ) + if self._closing or self._closed: + raise RuntimeError("cannot submit to a closing FeatureStore uploader") + while len(self._pending) >= self._settings.max_pending_steps: + self._condition.wait(self._settings.poll_interval_secs) + self._raise_if_failed_locked() + self._pending.append((global_step, table)) + self._condition.notify_all() + + def check_error(self) -> None: + """Surface a background failure at a safe training-thread boundary.""" + with self._condition: + self._raise_if_failed_locked() + + def close(self, raise_on_error: bool = True, drain: bool = True) -> None: + """Close the worker, draining only during a normal training shutdown.""" + with self._condition: + if self._closed: + if raise_on_error: + self._raise_if_failed_locked() + return + if not self._started: + self._closed = True + return + self._closing = True + self._aborting = not drain + self._condition.notify_all() + worker = self._worker + + if drain and worker is not None: + worker.join(timeout=self._settings.shutdown_timeout_secs) + if worker.is_alive(): + timeout_error = FeatureStoreUploadError( + "FeatureStore uploader did not drain before shutdown timeout" + ) + with self._condition: + if self._error is None: + self._error = timeout_error + self._aborting = True + self._condition.notify_all() + + with self._condition: + self._closed = True + if raise_on_error: + self._raise_if_failed_locked() + + def _run(self) -> None: + current_step: Optional[int] = None + try: + while True: + with self._condition: + if self._aborting: + return + if not self._pending: + if self._closing: + return + self._condition.wait(self._settings.poll_interval_secs) + continue + current_step, table = self._pending[0] + + self._upload_with_retries(current_step, table) + + with self._condition: + self._pending.popleft() + self._condition.notify_all() + current_step = None + except _UploadAborted: + return + except BaseException as exc: + step_context = ( + f" at global_step={current_step}" if current_step is not None else "" + ) + error = FeatureStoreUploadError( + f"FeatureStore delta upload failed{step_context}: {exc}" + ) + with self._condition: + if self._error is None: + self._error = error + self._condition.notify_all() + logger.error( + "FeatureStore delta upload failed%s: %s", + step_context, + exc, + exc_info=True, + ) + finally: + self._reset_view(suppress_errors=True) + + def _raise_if_failed_locked(self) -> None: + if self._error is not None: + raise self._error + + def _raise_if_aborting(self) -> None: + with self._condition: + if self._aborting: + raise _UploadAborted() + + def _upload_with_retries(self, global_step: int, table: pa.Table) -> None: + for attempt in range(1, self._settings.max_retries + 1): + self._raise_if_aborting() + try: + self._stream_upload(global_step, table) + return + except _UploadAborted: + raise + except BaseException as exc: + self._reset_view(suppress_errors=True) + if attempt >= self._settings.max_retries: + raise + logger.warning( + "FeatureStore delta upload attempt %s/%s failed for step %s " + "(%s); retrying after backoff", + attempt, + self._settings.max_retries, + global_step, + exc, + ) + if self._settings.retry_backoff_secs > 0: + time.sleep(self._settings.retry_backoff_secs * attempt) + raise AssertionError("unreachable FeatureStore retry state") + + def _allocate_timestamp_range(self, batch_count: int) -> Tuple[int, int]: + """Allocate a rank-locally monotonic timestamp range (in-memory only). + + Per-rank monotonicity is sufficient for Next-Ts incremental readers: + sharding is fixed for the lifetime of the process, so every key is + always republished by the same rank with a strictly newer timestamp. + """ + reserved = max(batch_count, 1) + range_start = max(int(self._clock_ms()), self._last_publish_ts + 1, 1) + range_end = range_start + reserved - 1 + self._last_publish_ts = range_end + return range_start, range_end + + def _stream_upload(self, global_step: int, table: pa.Table) -> None: + """Stream the in-memory delta table directly to the FeatureStore SDK.""" + view = self._get_view() + max_in_flight = int(getattr(view, "_max_workers", 1)) + + # Materialize the actual batch list so the ts range covers every batch + # the SDK will see. to_batches() splits each physical chunk independently, + # so ceil(total_rows / batch_size) undercounts multi-chunk (multi-FQN) + # tables and lets consecutive steps' timestamps collide. + batches = list(table.to_batches(max_chunksize=self._settings.upload_batch_size)) + total_batches = len(batches) or 1 + ts_range = self._allocate_timestamp_range(total_batches) + range_start = ts_range[0] + + completed_batches = 0 + window_batches = 0 + window_records = 0 + started_at = time.monotonic() + next_progress_batch = _FEATURE_STORE_PROGRESS_LOG_INTERVAL_BATCHES + logged_first_window = False + + logger.debug( + "FeatureStore delta upload started: step=%s rank=%s version=%s " + "batches=%s ts_range=%s-%s", + global_step, + self._rank, + self._settings.version, + total_batches, + ts_range[0], + ts_range[1], + ) + + try: + for batch in batches: + self._raise_if_aborting() + num_rows = self._submit_one_batch( + view, batch, range_start + completed_batches + ) + if num_rows == 0: + continue + completed_batches += 1 + window_batches += 1 + window_records += num_rows + + if window_batches < max_in_flight and completed_batches < total_batches: + continue + summary = view.write_flush() + self._validate_flush_summary( + summary, + expected_records=window_records, + expected_batches=window_batches, + ) + if ( + not logged_first_window + or completed_batches >= next_progress_batch + or completed_batches == total_batches + ): + log_progress = logger.info if logged_first_window else logger.debug + log_progress( + "FeatureStore delta upload progress: step=%s " + "batches=%s/%s elapsed_secs=%.1f", + global_step, + completed_batches, + total_batches, + time.monotonic() - started_at, + ) + logged_first_window = True + while next_progress_batch <= completed_batches: + next_progress_batch += ( + _FEATURE_STORE_PROGRESS_LOG_INTERVAL_BATCHES + ) + window_batches = 0 + window_records = 0 + + if window_batches > 0: + summary = view.write_flush() + self._validate_flush_summary( + summary, + expected_records=window_records, + expected_batches=window_batches, + ) + except BaseException: + try: + view.write_flush() + except BaseException: + pass + raise + + logger.info( + "FeatureStore delta upload completed: step=%s batches=%s elapsed_secs=%.1f", + global_step, + completed_batches, + time.monotonic() - started_at, + ) + + def _submit_one_batch(self, view: Any, batch: pa.RecordBatch, ts: int) -> int: + """Validate, build, and submit one batch; return submitted rows (0 skip). + + Dispatches on upload_format: ARROW streams a columnar RecordBatch through + write_features_arrow(); JSON falls back to the per-row write_features() + payload. Both reuse _validate_delta_batch so invariants are identical. + """ + if self._settings.upload_format == FEATURE_STORE_UPLOAD_FORMAT_DEFAULT: + wire_batch, num_rows = self._validate_and_build_arrow_batch(batch) + if num_rows == 0: + return 0 + view.write_features_arrow( + batch=wire_batch, + version=self._settings.version, + write_mode=FEATURE_STORE_WRITE_MODE, + ts=ts, + ) + else: + payload = self._validate_and_build_payload(batch) + num_rows = len(payload) + if num_rows == 0: + return 0 + view.write_features( + data=payload, + version=self._settings.version, + write_mode=FEATURE_STORE_WRITE_MODE, + ts=ts, + ) + return num_rows + + def _validate_delta_batch(self, batch: pa.RecordBatch) -> _DeltaBatch: + """Validate one delta batch and decode it for payload construction. + + Enforces the table_fqn / key_id / embedding / dimension invariants shared + by the JSON and Arrow upload paths. PK values are pre-remapped here + (INPUT_TILE=3 user-side keys -> non-user twin keys); dimension checks key + against the raw table_fqn, which is how the model contract is indexed. + + Args: + batch: One delta RecordBatch from the in-memory delta table. + + Returns: + Decoded _DeltaBatch; num_rows==0 when the batch carries no rows. + + Raises: + ValueError: On empty table_fqn, contract mismatch, reserved + key_id=-1, NaN/Inf embeddings, or dimension mismatch. + """ + num_rows = batch.num_rows + if num_rows == 0: + return _DeltaBatch( + num_rows=0, + remapped_fqns=[], + key_ids=np.empty(0, dtype=np.int64), + embedding_column=batch.column("embedding"), + flat_embeddings=np.empty(0, dtype=np.float32), + offsets=np.array([0], dtype=np.int32), + ) + table_fqns = batch.column("table_fqn").to_pylist() + for table_fqn in set(table_fqns): + if not table_fqn: + raise ValueError("delta shard table_fqn must not be empty") + if table_fqn not in self._embedding_dimensions: + raise ValueError( + "delta shard table_fqn is absent from model contract: " + f"{table_fqn!r}" + ) + + key_ids = batch.column("key_id").to_numpy(zero_copy_only=False) + if bool((key_ids == SPARSE_EMBEDDING_INVALID_KEY).any()): + raise ValueError( + "delta shard key_id=-1 is reserved as the Processor/" + "NvEmbeddings invalid-key sentinel" + ) + + embedding_column = cast(pa.ListArray, batch.column("embedding")) + offsets = embedding_column.offsets.to_numpy() + flat_embeddings = embedding_column.values.to_numpy(zero_copy_only=False) + # A sliced ListArray shares the whole chunk's child buffer, so scanning + # flat_embeddings in full would re-scan the chunk on every batch. Bound + # the NaN/Inf check to this batch's value range [offsets[0], offsets[-1]). + value_start = int(offsets[0]) + value_end = int(offsets[-1]) + if not bool(np.isfinite(flat_embeddings[value_start:value_end]).all()): + raise ValueError("delta embedding contains NaN or Inf") + lengths = np.diff(offsets) + expected_dims = np.array( + [self._embedding_dimensions[fqn] for fqn in table_fqns], + dtype=lengths.dtype, + ) + bad_rows = np.flatnonzero(lengths != expected_dims) + if bad_rows.size > 0: + row = int(bad_rows[0]) + raise ValueError( + f"delta embedding dimension mismatch for {table_fqns[row]!r}: " + f"expected={int(expected_dims[row])}, " + f"actual={int(lengths[row])}" + ) + + remap_cache = {fqn: remap_input_tile_user_key(fqn) for fqn in set(table_fqns)} + remapped_fqns = [remap_cache[fqn] for fqn in table_fqns] + return _DeltaBatch( + num_rows=num_rows, + remapped_fqns=remapped_fqns, + key_ids=key_ids, + embedding_column=embedding_column, + flat_embeddings=flat_embeddings, + offsets=offsets, + ) + + def _validate_and_build_payload( + self, + batch: pa.RecordBatch, + ) -> List[Dict[str, Any]]: + """Validate one delta batch and build the JSON SDK payload.""" + delta = self._validate_delta_batch(batch) + if delta.num_rows == 0: + return [] + return [ + { + FEATURE_STORE_PK_FIELD: delta.remapped_fqns[i], + FEATURE_STORE_SK_FIELD: int(delta.key_ids[i]), + FEATURE_STORE_VALUE_FIELD: delta.flat_embeddings[ + int(delta.offsets[i]) : int(delta.offsets[i + 1]) + ].copy(), + } + for i in range(delta.num_rows) + ] + + def _validate_and_build_arrow_batch( + self, + batch: pa.RecordBatch, + ) -> Tuple[Optional[pa.RecordBatch], int]: + """Validate one delta batch and build the Arrow IPC wire batch. + + Returns a RecordBatch with the configured PK/SK/embedding field names so + the SDK remaps them to its wire (pk/sk/embedding) columns. The embedding + column is reused zero-copy; only the string PK column and the int64 SK + column are rebuilt, avoiding the JSON path's per-row embedding deep-copy. + """ + delta = self._validate_delta_batch(batch) + if delta.num_rows == 0: + return None, 0 + pk_column = pa.array(delta.remapped_fqns, type=pa.string()) + sk_column = pa.array(delta.key_ids, type=pa.int64()) + wire_batch = pa.RecordBatch.from_arrays( + [pk_column, sk_column, delta.embedding_column], + names=[ + FEATURE_STORE_PK_FIELD, + FEATURE_STORE_SK_FIELD, + FEATURE_STORE_VALUE_FIELD, + ], + ) + return wire_batch, delta.num_rows + + @staticmethod + def _create_credentials_client() -> Any: + """Create the Alibaba Cloud credential provider (default chain).""" + try: + from alibabacloud_credentials.client import Client as CredClient + except ImportError as exc: + raise RuntimeError( + "alibabacloud_credentials is required when feature_store_config " + "is set; install it via: pip install alibabacloud_credentials" + ) from exc + return CredClient() + + def _create_client(self) -> Any: + """Construct a FeatureStoreClient with refreshed credentials. + + Single seam for credential resolution and client construction; tests + patch this method to inject a fake client. + """ + try: + from feature_store_py import FeatureStoreClient + except ImportError as exc: + raise RuntimeError( + "feature_store_py is required when feature_store_config is set" + ) from exc + credential = self._credentials_client.get_credential() + return FeatureStoreClient( + access_key_id=credential.access_key_id, + access_key_secret=credential.access_key_secret, + region=self._settings.region or None, + endpoint=self._settings.endpoint or None, + security_token=credential.security_token or None, + featuredb_username=os.environ.get("FEATUREDB_USERNAME") or None, + featuredb_password=os.environ.get("FEATUREDB_PASSWORD") or None, + ) + + def _get_view(self) -> Any: + if self._view is not None: + return self._view + client = self._create_client() + project = client.get_project(self._settings.project_name) + if project is None: + raise RuntimeError("configured FeatureStore project was not found") + view = self._get_or_create_view(project) + self._view = view + actual_fields = (view.pk_field, view.sk_field, view.embedding_field) + expected_fields = ( + FEATURE_STORE_PK_FIELD, + FEATURE_STORE_SK_FIELD, + FEATURE_STORE_VALUE_FIELD, + ) + if actual_fields != expected_fields: + raise RuntimeError( + "DynamicEmbedding FeatureView schema mismatch: " + f"expected={expected_fields}, actual={actual_fields}" + ) + sdk_batch_size = getattr(view, "_batch_size", FEATURE_STORE_SDK_BATCH_SIZE) + if ( + type(sdk_batch_size) is not int + or sdk_batch_size < self._settings.upload_batch_size + ): + raise RuntimeError( + "FeatureStore SDK batch_size is smaller than the configured outer " + "batch; one publish timestamp could span multiple HTTP requests" + ) + sdk_max_workers = getattr(view, "_max_workers", 1) + if type(sdk_max_workers) is not int or sdk_max_workers <= 0: + raise RuntimeError("FeatureStore SDK max_workers must be a positive int") + return view + + def _reset_view(self, suppress_errors: bool = False) -> None: + view = self._view + self._view = None + if view is not None: + try: + view.close(wait=True) + except BaseException as exc: + close_error = FeatureStoreUploadError( + "FeatureStore SDK writer close failed" + ) + if not suppress_errors: + raise close_error from exc + with self._condition: + if self._error is None: + self._error = close_error + self._condition.notify_all() + logger.error( + "Failed to close FeatureStore SDK writer cleanly (%s)", + type(exc).__name__, + ) + + def _get_or_create_view(self, project: Any) -> Any: + """Return the configured DynamicEmbedding view, creating it if absent. + + Only the primary (rank-zero) uploader creates the view; other ranks open + a handle to the view that the primary published before they started. + Schema compatibility is checked once on the data-plane writer in + ``_get_view``; creation-time provisioning (TTL/shard/replication) does + not affect upload compatibility and is not re-validated here. + """ + if not self._manage_remote_view: + view = project.get_dynamic_embedding_feature_view( + self._settings.feature_view_name + ) + if view is None: + view = self._wait_for_dynamic_embedding_view(project) + if view is None: + raise RuntimeError( + "configured DynamicEmbedding FeatureView was not found; " + "the rank-zero uploader must create it before other ranks " + "start" + ) + self._view = view + return view + provisioned = False + view = project.get_dynamic_embedding_feature_view( + self._settings.feature_view_name + ) + if view is None: + create_error: Optional[Exception] = None + try: + entity_name = self._get_or_create_entity(project) + view = project.create_dynamic_embedding_feature_view( + name=self._settings.feature_view_name, + entity=entity_name, + pk_field_name=FEATURE_STORE_PK_FIELD, + sk_field_name=FEATURE_STORE_SK_FIELD, + embedding_field_name=FEATURE_STORE_VALUE_FIELD, + pk_field_type="STRING", + sk_field_type="INT64", + ttl=self._settings.feature_view_ttl_secs, + shard_count=self._settings.feature_view_shard_count, + replication_count=self._settings.feature_view_replication_count, + ) + provisioned = True + except Exception as exc: + create_error = exc + view = self._wait_for_dynamic_embedding_view(project) + if view is None: + error = RuntimeError( + "failed to create configured DynamicEmbedding FeatureView" + ) + if create_error is not None: + raise error from create_error + raise error + self._view = view + if provisioned: + logger.info( + "Created DynamicEmbedding FeatureView: project=%s entity=%s view=%s", + self._settings.project_name, + FEATURE_STORE_DEFAULT_ENTITY_NAME, + self._settings.feature_view_name, + ) + return view + + def _get_or_create_entity(self, project: Any) -> str: + """Return the default DynamicEmbedding entity name, creating it if absent. + + Only the rank-zero uploader runs this, immediately before it creates the + view, so there is no cross-rank race; a concurrent external creator is + handled by re-getting the entity if the create call fails. + """ + if project.get_entity(FEATURE_STORE_DEFAULT_ENTITY_NAME) is not None: + return FEATURE_STORE_DEFAULT_ENTITY_NAME + try: + project.create_entity( + FEATURE_STORE_DEFAULT_ENTITY_NAME, + FEATURE_STORE_DEFAULT_ENTITY_JOIN_ID, + ) + except Exception as exc: + if project.get_entity(FEATURE_STORE_DEFAULT_ENTITY_NAME) is None: + raise RuntimeError( + "failed to create default DynamicEmbedding entity " + f"{FEATURE_STORE_DEFAULT_ENTITY_NAME!r}" + ) from exc + logger.info( + "Created FeatureStore entity: project=%s entity=%s", + self._settings.project_name, + FEATURE_STORE_DEFAULT_ENTITY_NAME, + ) + return FEATURE_STORE_DEFAULT_ENTITY_NAME + + def _wait_for_dynamic_embedding_view(self, project: Any) -> Any: + """Bounded re-get after a concurrent or partially completed create.""" + last_error: Optional[Exception] = None + for attempt in range(1, self._settings.max_retries + 1): + try: + view = project.get_dynamic_embedding_feature_view( + self._settings.feature_view_name + ) + except Exception as exc: + last_error = exc + else: + if view is not None: + return view + if ( + attempt < self._settings.max_retries + and self._settings.retry_backoff_secs > 0 + ): + time.sleep(self._settings.retry_backoff_secs * attempt) + if last_error is not None: + raise RuntimeError( + "DynamicEmbedding FeatureView did not become ready after creation" + ) from last_error + return None + + @staticmethod + def _validate_flush_summary( + summary: Any, expected_records: int, expected_batches: int + ) -> None: + required = { + "total_batches", + "failed_batches", + "total_records", + "success_records", + "failed_records", + } + if not isinstance(summary, dict) or not required.issubset(summary): + raise RuntimeError("FeatureStore write_flush returned an invalid summary") + if ( + int(summary["total_batches"]) != expected_batches + or int(summary["failed_batches"]) != 0 + or int(summary["failed_records"]) != 0 + or int(summary["success_records"]) != int(summary["total_records"]) + or int(summary["total_records"]) != expected_records + ): + raise RuntimeError("FeatureStore write_flush reported incomplete writes") diff --git a/tzrec/utils/feature_store_delta_uploader_test.py b/tzrec/utils/feature_store_delta_uploader_test.py new file mode 100644 index 00000000..d19492eb --- /dev/null +++ b/tzrec/utils/feature_store_delta_uploader_test.py @@ -0,0 +1,1077 @@ +# Copyright (c) 2025, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import sys +import threading +import types +import unittest +from unittest import mock + +import numpy as np +import pyarrow as pa +from google.protobuf.descriptor import FieldDescriptor + +from tzrec.protos.train_pb2 import FeatureStoreConfig +from tzrec.utils.delta_embedding_dump import _DELTA_DUMP_SCHEMA +from tzrec.utils.feature_store_delta_uploader import ( + FEATURE_STORE_DEFAULT_ENTITY_JOIN_ID, + FEATURE_STORE_DEFAULT_ENTITY_NAME, + FeatureStoreDeltaUploader, + FeatureStoreUploadError, + FeatureStoreUploadSettings, +) + + +def _feature_store_config(**overrides) -> FeatureStoreConfig: + config = FeatureStoreConfig( + region="cn-test", + project_name="project_a", + feature_view_name="shared_embeddings", + version="model_a@export_1", + upload_batch_size=2, + max_retries=1, + retry_backoff_secs=0, + shutdown_timeout_secs=5, + max_pending_steps=8, + poll_interval_secs=1, + ) + for name, value in overrides.items(): + setattr(config, name, value) + return config + + +def _row( + step: int, + rank: int, + key_id: int, + values, + name: str = "user_emb", + world_size: int = 1, +): + return { + "global_step": step, + "rank": rank, + "world_size": world_size, + "feature_name": "user_id", + "table_fqn": f"model.ebc.embedding_bags.{name}", + "key_id": key_id, + "embedding": values, + "source": "model_delta_tracker", + } + + +def _delta_table(rows) -> pa.Table: + if rows: + return pa.Table.from_pylist(rows, schema=_DELTA_DUMP_SCHEMA) + return _DELTA_DUMP_SCHEMA.empty_table() + + +class _FakeView: + pk_field = "embedding_name" + sk_field = "key_id" + embedding_field = "embedding" + + def __init__(self, summaries=None, close_error=None, max_workers=4): + self.calls = [] + self.arrow_calls = [] + self.closed = [] + self.flush_calls = [] + self._summaries = list(summaries or []) + self._close_error = close_error + self._batch_size = 1000 + self._max_workers = max_workers + self._pending_sizes = [] + + def write_features(self, **kwargs): + self.calls.append(kwargs) + self._pending_sizes.append(len(kwargs["data"])) + + def write_features_arrow(self, *, batch, version, write_mode, ts): + # Decode the Arrow wire batch into the same {data, version, write_mode, + # ts} call shape as write_features, so the existing JSON-path assertions + # also exercise the default Arrow path unchanged. The raw batch is kept + # on arrow_calls for column-type / column-name assertions. + self.arrow_calls.append( + {"batch": batch, "version": version, "write_mode": write_mode, "ts": ts} + ) + data = [ + { + self.pk_field: pk, + self.sk_field: int(sk), + self.embedding_field: np.asarray(emb, dtype=np.float32), + } + for pk, sk, emb in zip( + batch.column(self.pk_field).to_pylist(), + batch.column(self.sk_field).to_pylist(), + batch.column(self.embedding_field).to_pylist(), + ) + ] + self.calls.append( + {"data": data, "version": version, "write_mode": write_mode, "ts": ts} + ) + self._pending_sizes.append(len(data)) + + def write_flush(self): + pending_sizes = self._pending_sizes + self._pending_sizes = [] + self.flush_calls.append(pending_sizes) + if self._summaries: + return self._summaries.pop(0) + total_records = sum(pending_sizes) + return { + "total_batches": len(pending_sizes), + "failed_batches": 0, + "total_records": total_records, + "success_records": total_records, + "failed_records": 0, + "errors": [], + } + + def close(self, wait=True): + self.closed.append(wait) + if self._close_error is not None: + raise self._close_error + + +class _BlockingView(_FakeView): + def __init__(self): + super().__init__() + self.flush_started = threading.Event() + self.release_flush = threading.Event() + self.close_finished = threading.Event() + + def write_flush(self): + self.flush_started.set() + self.release_flush.wait(timeout=5) + return super().write_flush() + + def close(self, wait=True): + super().close(wait=wait) + self.close_finished.set() + + +class _FakeEntity: + def __init__(self, name): + self.feature_entity_name = name + + +class _FakeProject: + def __init__( + self, + view, + *, + created_view=None, + create_error=None, + view_after_create_error=None, + entity="existing_entity", + entity_create_error=None, + ): + self._view = view + self._created_view = created_view + self._create_error = create_error + self._view_after_create_error = view_after_create_error + self._entity = entity + self._entity_create_error = entity_create_error + self.dynamic_get_calls = [] + self.create_calls = [] + self.entity_get_calls = [] + self.entity_create_calls = [] + + def get_dynamic_embedding_feature_view(self, name): + self.dynamic_get_calls.append(name) + return self._view + + def create_dynamic_embedding_feature_view(self, **kwargs): + self.create_calls.append(kwargs) + if self._create_error is not None: + self._view = self._view_after_create_error + raise self._create_error + self._view = self._created_view or _FakeView() + return self._view + + def get_entity(self, name): + self.entity_get_calls.append(name) + return self._entity + + def create_entity(self, name, join_id, parent_feature_entity_name=None): + self.entity_create_calls.append((name, join_id)) + if self._entity_create_error is not None: + raise self._entity_create_error + self._entity = _FakeEntity(name) + return self._entity + + +class _FakeCredential: + access_key_id = "fake-ak" + access_key_secret = "fake-sk" + security_token = "fake-sts" + + +class _FakeCredentialsClient: + def get_credential(self): + return _FakeCredential() + + +class _FakeClient: + def __init__(self, project, kwargs): + self._project = project + self.kwargs = kwargs + + def get_project(self, name): + return self._project + + +class _FakeClientFactory: + def __init__(self, view, **project_kwargs): + self.view = view + self.calls = [] + self.project = _FakeProject(view, **project_kwargs) + + def __call__(self, **kwargs): + self.calls.append(kwargs) + return _FakeClient(self.project, kwargs) + + +class _SequencedClientFactory: + def __init__(self, views): + self._projects = [_FakeProject(view) for view in views] + self.calls = [] + + def __call__(self, **kwargs): + self.calls.append(kwargs) + return _FakeClient(self._projects.pop(0), kwargs) + + +class FeatureStoreDeltaUploaderTest(unittest.TestCase): + def setUp(self): + self._cred_patch = mock.patch.object( + FeatureStoreDeltaUploader, + "_create_credentials_client", + return_value=_FakeCredentialsClient(), + ) + self._cred_patch.start() + self.addCleanup(self._cred_patch.stop) + + def _uploader(self, config=None, **kwargs): + client_factory = kwargs.pop("client_factory", None) + kwargs.setdefault( + "embedding_dimensions", {"model.ebc.embedding_bags.user_emb": 2} + ) + uploader = FeatureStoreDeltaUploader( + config or _feature_store_config(), **kwargs + ) + if client_factory is not None: + # The production ctor no longer accepts a client factory; install the + # test's fake on the _create_client instance seam so each uploader + # keeps its own fake project/view wiring and construction counts. + uploader._create_client = lambda *a, **k: client_factory() + return uploader + + def test_proto_groups_required_fields_before_optional_fields(self): + required_fields = [ + "region", + "project_name", + "feature_view_name", + "version", + ] + optional_fields = [ + "endpoint", + "upload_batch_size", + "max_retries", + "retry_backoff_secs", + "shutdown_timeout_secs", + "max_pending_steps", + "poll_interval_secs", + "feature_view_ttl_secs", + "feature_view_shard_count", + "feature_view_replication_count", + "retain_local_dump", + "upload_format", + ] + fields = list(FeatureStoreConfig.DESCRIPTOR.fields) + + self.assertEqual( + [field.name for field in fields], required_fields + optional_fields + ) + self.assertTrue( + all( + field.label == FieldDescriptor.LABEL_REQUIRED + for field in fields[: len(required_fields)] + ) + ) + self.assertTrue( + all( + field.label == FieldDescriptor.LABEL_OPTIONAL + for field in fields[len(required_fields) :] + ) + ) + self.assertEqual( + [field.number for field in fields], + [1, 2, 3] + list(range(5, 16)) + [17, 18], + ) + for field_name in required_fields: + with self.subTest(field_name=field_name): + config = _feature_store_config() + config.ClearField(field_name) + self.assertFalse(config.IsInitialized()) + self.assertIn(field_name, config.FindInitializationErrors()) + with self.assertRaisesRegex(ValueError, field_name): + FeatureStoreUploadSettings.from_proto(config) + + def test_version_is_required_and_must_be_explicit(self): + config = _feature_store_config() + config.ClearField("version") + + self.assertFalse(config.IsInitialized()) + self.assertIn("version", config.FindInitializationErrors()) + with self.assertRaisesRegex(ValueError, "required fields.*version"): + FeatureStoreUploadSettings.from_proto(config) + + with self.assertRaisesRegex(ValueError, "explicit non-default version"): + FeatureStoreUploadSettings.from_proto( + _feature_store_config(version="default") + ) + + def test_region_fallback_and_config_validation(self): + config = _feature_store_config(region="") + with mock.patch.dict(os.environ, {"ALIBABA_CLOUD_REGION": "cn-env"}): + settings = FeatureStoreUploadSettings.from_proto(config) + self.assertEqual(settings.region, "cn-env") + + with self.assertRaisesRegex(ValueError, "must be <= 1000"): + FeatureStoreUploadSettings.from_proto( + _feature_store_config(upload_batch_size=1001) + ) + with self.assertRaisesRegex(ValueError, "shard_count must be in"): + FeatureStoreUploadSettings.from_proto( + _feature_store_config(feature_view_shard_count=21) + ) + with self.assertRaisesRegex(ValueError, "replication_count must be in"): + FeatureStoreUploadSettings.from_proto( + _feature_store_config(feature_view_replication_count=4) + ) + + def test_start_reuses_existing_dynamic_embedding_feature_view(self): + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual(factory.project.dynamic_get_calls, ["shared_embeddings"]) + self.assertEqual(factory.project.create_calls, []) + self.assertEqual(view.closed, [True]) + + def test_create_client_forwards_only_credential_kwargs(self): + """_create_client forwards the fixed credential allowlist, no extras.""" + recorded = {} + + class _RecordingClient: + def __init__(self, **kwargs): + recorded.update(kwargs) + + fake_module = types.ModuleType("feature_store_py") + fake_module.FeatureStoreClient = _RecordingClient + with mock.patch.dict(sys.modules, {"feature_store_py": fake_module}): + uploader = self._uploader() + client = uploader._create_client() + self.assertIsInstance(client, _RecordingClient) + self.assertEqual( + set(recorded), + { + "access_key_id", + "access_key_secret", + "region", + "endpoint", + "security_token", + "featuredb_username", + "featuredb_password", + }, + ) + self.assertNotIn("test_mode", recorded) + + def test_start_creates_missing_dynamic_embedding_feature_view(self): + created_view = _FakeView() + factory = _FakeClientFactory(None, created_view=created_view) + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual(factory.project.dynamic_get_calls, ["shared_embeddings"]) + self.assertEqual( + factory.project.create_calls, + [ + { + "name": "shared_embeddings", + "entity": FEATURE_STORE_DEFAULT_ENTITY_NAME, + "pk_field_name": "embedding_name", + "sk_field_name": "key_id", + "embedding_field_name": "embedding", + "pk_field_type": "STRING", + "sk_field_type": "INT64", + "ttl": 1296000, + "shard_count": 20, + "replication_count": 1, + } + ], + ) + self.assertEqual(factory.project.entity_create_calls, []) + self.assertEqual(created_view.closed, [True]) + + def test_start_recovers_from_concurrent_feature_view_creation(self): + concurrent_view = _FakeView() + factory = _FakeClientFactory( + None, + create_error=RuntimeError("already exists"), + view_after_create_error=concurrent_view, + ) + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual(len(factory.project.create_calls), 1) + self.assertEqual( + factory.project.dynamic_get_calls, + ["shared_embeddings", "shared_embeddings"], + ) + self.assertEqual(concurrent_view.closed, [True]) + + def test_start_closes_new_feature_view_with_incompatible_schema(self): + created_view = _FakeView() + created_view.pk_field = "wrong_pk" + factory = _FakeClientFactory(None, created_view=created_view) + uploader = self._uploader(client_factory=factory) + + with self.assertRaisesRegex(RuntimeError, "schema mismatch"): + uploader.start() + + self.assertEqual(len(factory.project.create_calls), 1) + self.assertEqual(created_view.closed, [True]) + + def test_start_creates_missing_view_without_version_precheck(self): + created_view = _FakeView() + factory = _FakeClientFactory(None, created_view=created_view) + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual(len(factory.project.create_calls), 1) + self.assertEqual(created_view.closed, [True]) + + def test_start_raises_when_view_creation_fails_and_view_never_appears(self): + factory = _FakeClientFactory(None, create_error=ValueError("boom")) + uploader = self._uploader(client_factory=factory) + + with self.assertRaisesRegex( + RuntimeError, "failed to create configured DynamicEmbedding FeatureView" + ): + uploader.start() + + self.assertEqual(len(factory.project.create_calls), 1) + self.assertEqual(factory.project.entity_create_calls, []) + + def test_start_creates_default_entity_when_it_does_not_exist(self): + created_view = _FakeView() + factory = _FakeClientFactory(None, created_view=created_view, entity=None) + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual( + factory.project.entity_get_calls, [FEATURE_STORE_DEFAULT_ENTITY_NAME] + ) + self.assertEqual( + factory.project.entity_create_calls, + [ + ( + FEATURE_STORE_DEFAULT_ENTITY_NAME, + FEATURE_STORE_DEFAULT_ENTITY_JOIN_ID, + ) + ], + ) + self.assertEqual( + factory.project.create_calls[0]["entity"], + FEATURE_STORE_DEFAULT_ENTITY_NAME, + ) + self.assertEqual(created_view.closed, [True]) + + def test_start_recovers_from_concurrent_default_entity_creation(self): + created_view = _FakeView() + factory = _FakeClientFactory( + None, + created_view=created_view, + entity=None, + entity_create_error=RuntimeError("entity already exists"), + ) + # A concurrent creator wins the race: the first get_entity sees no entity, + # create_entity then fails, and the entity is visible on the retry. + entity_results = [None, _FakeEntity(FEATURE_STORE_DEFAULT_ENTITY_NAME)] + + def get_entity(name): + factory.project.entity_get_calls.append(name) + return entity_results.pop(0) + + factory.project.get_entity = get_entity + uploader = self._uploader(client_factory=factory) + + uploader.start() + uploader.close() + + self.assertEqual( + factory.project.entity_create_calls, + [ + ( + FEATURE_STORE_DEFAULT_ENTITY_NAME, + FEATURE_STORE_DEFAULT_ENTITY_JOIN_ID, + ) + ], + ) + self.assertEqual(len(factory.project.create_calls), 1) + self.assertEqual(created_view.closed, [True]) + + def test_non_primary_uploader_opens_view_without_create_or_metadata_checks(self): + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + rank=1, + manage_remote_view=False, + client_factory=factory, + clock_ms=lambda: 100, + ) + + uploader.start() + uploader.submit(10, _delta_table([_row(10, 1, 7, [1.0, 2.0])])) + uploader.close() + + self.assertEqual(factory.project.dynamic_get_calls, ["shared_embeddings"]) + self.assertEqual(factory.project.create_calls, []) + self.assertEqual(len(view.calls), 1) + self.assertEqual(view.calls[0]["data"][0]["key_id"], 7) + + def test_non_primary_uploader_fails_when_view_is_missing(self): + factory = _FakeClientFactory(None) + uploader = self._uploader( + rank=1, + manage_remote_view=False, + client_factory=factory, + ) + + with self.assertRaisesRegex(RuntimeError, "rank-zero uploader must create"): + uploader.start() + + self.assertEqual(factory.project.create_calls, []) + + def test_submit_requires_started_uploader(self): + uploader = self._uploader(client_factory=_FakeClientFactory(_FakeView())) + with self.assertRaisesRegex(RuntimeError, "start.*before submit"): + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0])])) + uploader.close() + + def test_complete_step_uploads_merge_with_stable_version_and_ts(self): + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + client_factory=factory, + clock_ms=lambda: 123456, + ) + uploader.start() + uploader.submit( + 10, + _delta_table( + [ + _row(10, 0, 1, [1.0, 2.0]), + _row(10, 0, 2, [3.0, 4.0]), + _row(10, 0, 3, [0.0, 0.0]), + ] + ), + ) + uploader.close() + + self.assertEqual(len(view.calls), 2) + self.assertEqual([len(call["data"]) for call in view.calls], [2, 1]) + self.assertEqual(view.flush_calls, [[2, 1]]) + self.assertEqual({call["version"] for call in view.calls}, {"model_a@export_1"}) + self.assertEqual({call["write_mode"] for call in view.calls}, {"MERGE"}) + self.assertEqual([call["ts"] for call in view.calls], [123456, 123457]) + self.assertEqual(view.calls[1]["data"][0]["embedding"].tolist(), [0.0, 0.0]) + self.assertEqual(view.closed, [True]) + + def test_upload_uses_bounded_sdk_worker_windows(self): + view = _FakeView(max_workers=2) + uploader = self._uploader( + _feature_store_config(upload_batch_size=1), + client_factory=_FakeClientFactory(view), + clock_ms=lambda: 100, + ) + + uploader.start() + uploader.submit( + 10, _delta_table([_row(10, 0, key, [1.0, 2.0]) for key in range(1, 6)]) + ) + uploader.close() + + self.assertEqual([call["ts"] for call in view.calls], [100, 101, 102, 103, 104]) + self.assertEqual(view.flush_calls, [[1, 1], [1, 1], [1]]) + + def test_first_positive_dump_step_is_not_filtered(self): + view = _FakeView() + uploader = self._uploader( + client_factory=_FakeClientFactory(view), + clock_ms=lambda: 100, + ) + + uploader.start() + uploader.submit(1, _delta_table([_row(1, 0, 1, [1.0, 2.0])])) + uploader.close() + + self.assertEqual(len(view.calls), 1) + self.assertEqual(view.calls[0]["ts"], 100) + self.assertEqual(view.calls[0]["version"], "model_a@export_1") + + def test_submit_rejects_step_zero(self): + uploader = self._uploader(client_factory=_FakeClientFactory(_FakeView())) + + uploader.start() + try: + with self.assertRaisesRegex(ValueError, "global_step must be > 0"): + uploader.submit(0, _delta_table([])) + finally: + uploader.close() + + def test_flush_failure_raises_error(self): + failed_summary = { + "total_batches": 2, + "failed_batches": 1, + "total_records": 3, + "success_records": 2, + "failed_records": 1, + "errors": ["failed future"], + } + view = _FakeView([failed_summary]) + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit( + 10, + _delta_table( + [ + _row(10, 0, 1, [1.0, 2.0]), + _row(10, 0, 2, [3.0, 4.0]), + _row(10, 0, 3, [5.0, 6.0]), + ] + ), + ) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + self.assertEqual(view.flush_calls, [[2, 1], []]) + + def test_retry_uses_fresh_view_and_newer_timestamp_range(self): + failed_summary = { + "total_batches": 1, + "failed_batches": 1, + "total_records": 1, + "success_records": 0, + "failed_records": 1, + "errors": ["failed future"], + } + first_view = _FakeView([failed_summary]) + second_view = _FakeView() + factory = _SequencedClientFactory([first_view, second_view]) + uploader = self._uploader( + _feature_store_config(max_retries=2), + client_factory=factory, + clock_ms=lambda: 777, + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0])])) + uploader.close() + + self.assertEqual(len(factory.calls), 2) + self.assertEqual(first_view.closed, [True]) + self.assertEqual(second_view.closed, [True]) + all_calls = first_view.calls + second_view.calls + self.assertEqual({call["version"] for call in all_calls}, {"model_a@export_1"}) + self.assertEqual([call["ts"] for call in all_calls], [777, 778]) + + def test_merge_does_not_require_preprovisioned_version(self): + view = _FakeView() + uploader = self._uploader(client_factory=_FakeClientFactory(view)) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0])])) + uploader.close() + + self.assertEqual(len(view.calls), 1) + self.assertEqual(view.calls[0]["write_mode"], "MERGE") + + def test_non_draining_close_stops_without_commit(self): + view = _BlockingView() + uploader = self._uploader( + _feature_store_config(upload_batch_size=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit( + 10, _delta_table([_row(10, 0, key, [1.0, 2.0]) for key in range(1, 10)]) + ) + self.assertTrue(view.flush_started.wait(timeout=5)) + uploader.close(raise_on_error=False, drain=False) + view.release_flush.set() + self.assertTrue(view.close_finished.wait(timeout=5)) + self.assertTrue(len(view.calls) < 9) + + def test_signed_int64_key_is_preserved(self): + large_key = (1 << 63) - 1 + view = _FakeView() + uploader = self._uploader(client_factory=_FakeClientFactory(view)) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, large_key, [1.0, 2.0])])) + uploader.close() + + self.assertEqual(view.calls[0]["data"][0]["key_id"], large_key) + + def test_reserved_invalid_key_is_rejected(self): + view = _FakeView() + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, -1, [1.0, 2.0])])) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + self.assertEqual(view.calls, []) + + def test_empty_table_upload_writes_nothing(self): + view = _FakeView() + uploader = self._uploader(client_factory=_FakeClientFactory(view)) + uploader.start() + uploader.submit(10, _delta_table([])) + uploader.close() + + self.assertEqual(view.calls, []) + + def test_dimension_and_finite_value_validation(self): + view = _FakeView() + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0, 3.0])])) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + self.assertEqual(view.calls, []) + + view = _FakeView() + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [float("nan"), 2.0])])) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + self.assertEqual(view.calls, []) + + # Inf is also rejected (np.isfinite covers both; a regression to + # np.isnan would let Inf embeddings through undetected). + view = _FakeView() + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [float("inf"), 2.0])])) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + self.assertEqual(view.calls, []) + + def test_in_memory_timestamp_monotonicity_across_steps(self): + view = _FakeView() + uploader = self._uploader( + client_factory=_FakeClientFactory(view), + clock_ms=lambda: 100, + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0])])) + uploader.submit(20, _delta_table([_row(20, 0, 2, [3.0, 4.0])])) + uploader.close() + + ts_values = [call["ts"] for call in view.calls] + self.assertEqual(ts_values, [100, 101]) + + def test_in_memory_timestamp_monotonicity_across_retries(self): + failed_summary = { + "total_batches": 1, + "failed_batches": 1, + "total_records": 1, + "success_records": 0, + "failed_records": 1, + } + first_view = _FakeView([failed_summary]) + second_view = _FakeView() + factory = _SequencedClientFactory([first_view, second_view]) + uploader = self._uploader( + _feature_store_config(max_retries=2), + client_factory=factory, + clock_ms=lambda: 500, + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0])])) + uploader.close() + + self.assertEqual([call["ts"] for call in first_view.calls], [500]) + self.assertEqual([call["ts"] for call in second_view.calls], [501]) + + def test_data_parallel_ranks_upload_duplicate_keys_independently(self): + rank0_view = _FakeView() + rank1_view = _FakeView() + rank0 = self._uploader( + client_factory=_FakeClientFactory(rank0_view), + clock_ms=lambda: 100, + ) + rank1 = self._uploader( + rank=1, + manage_remote_view=False, + client_factory=_FakeClientFactory(rank1_view), + clock_ms=lambda: 100, + ) + rank0.start() + rank1.start() + rank0.submit(10, _delta_table([_row(10, 0, 1, [1.0, 2.0], world_size=2)])) + rank1.submit(10, _delta_table([_row(10, 1, 1, [1.0, 2.0], world_size=2)])) + rank0.close() + rank1.close() + + rank0_keys = [ + item["key_id"] for call in rank0_view.calls for item in call["data"] + ] + rank1_keys = [ + item["key_id"] for call in rank1_view.calls for item in call["data"] + ] + self.assertEqual(rank0_keys, [1]) + self.assertEqual(rank1_keys, [1]) + self.assertEqual({call["write_mode"] for call in rank0_view.calls}, {"MERGE"}) + self.assertEqual({call["write_mode"] for call in rank1_view.calls}, {"MERGE"}) + + def test_close_error_surfaces_via_check_error(self): + view = _FakeView() + uploader = self._uploader( + _feature_store_config(max_retries=1), + client_factory=_FakeClientFactory(view), + ) + uploader.start() + uploader.submit(10, _delta_table([_row(10, 0, -1, [1.0, 2.0])])) + with self.assertRaises(FeatureStoreUploadError): + uploader.close() + with self.assertRaises(FeatureStoreUploadError): + uploader.check_error() + + def test_default_upload_format_is_arrow(self): + # An unset upload_format inherits the proto default ARROW. + settings = FeatureStoreUploadSettings.from_proto(_feature_store_config()) + self.assertEqual(settings.upload_format, "ARROW") + + def test_rejects_unknown_upload_format(self): + config = _feature_store_config() + config.upload_format = "protobuf" + with self.assertRaisesRegex(ValueError, "upload_format must be one of"): + FeatureStoreUploadSettings.from_proto(config) + + def test_json_path_routes_through_write_features(self): + # Explicit JSON keeps the legacy per-row write_features payload path. + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + _feature_store_config(upload_format="JSON"), + client_factory=factory, + clock_ms=lambda: 123456, + ) + uploader.start() + uploader.submit( + 10, + _delta_table( + [ + _row(10, 0, 1, [1.0, 2.0]), + _row(10, 0, 2, [3.0, 4.0]), + _row(10, 0, 3, [0.0, 0.0]), + ] + ), + ) + uploader.close() + + self.assertEqual(view.arrow_calls, []) + self.assertEqual(len(view.calls), 2) + self.assertEqual([len(call["data"]) for call in view.calls], [2, 1]) + self.assertEqual(view.flush_calls, [[2, 1]]) + self.assertEqual({call["version"] for call in view.calls}, {"model_a@export_1"}) + self.assertEqual({call["write_mode"] for call in view.calls}, {"MERGE"}) + self.assertEqual([call["ts"] for call in view.calls], [123456, 123457]) + self.assertEqual(view.calls[0]["data"][0]["key_id"], 1) + self.assertEqual( + view.calls[0]["data"][0]["embedding_name"], + "model.ebc.embedding_bags.user_emb", + ) + self.assertTrue( + np.array_equal( + view.calls[0]["data"][0]["embedding"], + np.array([1.0, 2.0], dtype=np.float32), + ) + ) + # The sliced second batch exercises the offset-indexed embedding slice. + self.assertTrue( + np.array_equal( + view.calls[1]["data"][0]["embedding"], + np.array([0.0, 0.0], dtype=np.float32), + ) + ) + self.assertEqual(view.closed, [True]) + + def test_arrow_path_builds_wire_batch_columns(self): + # Default ARROW path builds a wire RecordBatch whose configured field + # names the SDK remaps to its pk/sk/embedding wire columns. + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + _feature_store_config(upload_batch_size=2), + client_factory=factory, + clock_ms=lambda: 100, + ) + uploader.start() + uploader.submit( + 10, + _delta_table( + [ + _row(10, 0, 1, [1.0, 2.0]), + _row(10, 0, 2, [3.0, 4.0]), + _row(10, 0, 3, [5.0, 6.0]), + ] + ), + ) + uploader.close() + + self.assertEqual(len(view.arrow_calls), 2) + self.assertEqual([c["ts"] for c in view.arrow_calls], [100, 101]) + self.assertEqual({c["version"] for c in view.arrow_calls}, {"model_a@export_1"}) + self.assertEqual({c["write_mode"] for c in view.arrow_calls}, {"MERGE"}) + + batch0 = view.arrow_calls[0]["batch"] + self.assertEqual(batch0.schema.names, ["embedding_name", "key_id", "embedding"]) + self.assertEqual(batch0.num_rows, 2) + # PK is the remapped table_fqn (string); SK stays int64 (the SDK casts to + # the string wire type); embedding is list reused zero-copy. + self.assertEqual(batch0.column("embedding_name").type, pa.string()) + self.assertEqual(batch0.column("key_id").type, pa.int64()) + self.assertEqual(batch0.column("embedding").type, pa.list_(pa.float32())) + self.assertEqual( + batch0.column("embedding_name").to_pylist(), + ["model.ebc.embedding_bags.user_emb"] * 2, + ) + self.assertEqual(batch0.column("key_id").to_pylist(), [1, 2]) + self.assertEqual( + batch0.column("embedding").to_pylist(), + [[1.0, 2.0], [3.0, 4.0]], + ) + self.assertEqual(view.closed, [True]) + + def test_multi_chunk_table_keeps_timestamps_monotonic_across_steps(self): + # A multi-FQN delta table concatenates one chunk per FQN; to_batches() + # splits each chunk independently, so the ts range must cover every + # actual batch or a stuck clock reuses a prior step's timestamps and + # Next-Ts incremental readers miss updates. + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + _feature_store_config(upload_batch_size=1000), + client_factory=factory, + clock_ms=lambda: 100, + embedding_dimensions={ + "model.ebc.embedding_bags.user_emb": 2, + "model.ebc.embedding_bags.item_emb": 2, + }, + ) + uploader.start() + rows_a = [_row(10, 0, k, [1.0, 2.0]) for k in range(5)] + rows_b = [_row(10, 0, k, [3.0, 4.0], name="item_emb") for k in range(5)] + table = pa.concat_tables([_delta_table(rows_a), _delta_table(rows_b)]) + uploader.submit(10, table) + uploader.submit(20, table) + uploader.close() + + ts_values = [call["ts"] for call in view.calls] + # Two chunks -> two batches per step; the stuck clock forces step 2 to + # start strictly after step 1's last ts (101), i.e. 102. + self.assertEqual(ts_values, [100, 101, 102, 103]) + self.assertEqual([len(call["data"]) for call in view.calls], [5, 5, 5, 5]) + self.assertEqual(view.closed, [True]) + + def test_mixed_fqn_dimensions_validated_per_row(self): + # _validate_delta_batch keys the expected dimension per row, so a batch + # carrying multiple FQNs with different dimensions must not be + # mis-flagged as a dimension mismatch. + view = _FakeView() + factory = _FakeClientFactory(view) + uploader = self._uploader( + _feature_store_config(upload_batch_size=1000), + client_factory=factory, + clock_ms=lambda: 100, + embedding_dimensions={ + "model.ebc.embedding_bags.user_emb": 2, + "model.ebc.embedding_bags.item_emb": 3, + }, + ) + uploader.start() + uploader.submit( + 10, + _delta_table( + [ + _row(10, 0, 1, [1.0, 2.0]), + _row(10, 0, 2, [3.0, 4.0, 5.0], name="item_emb"), + ] + ), + ) + uploader.close() + + self.assertEqual(len(view.calls), 1) + self.assertEqual(len(view.calls[0]["data"]), 2) + self.assertEqual( + view.calls[0]["data"][0]["embedding_name"], + "model.ebc.embedding_bags.user_emb", + ) + self.assertEqual( + view.calls[0]["data"][1]["embedding_name"], + "model.ebc.embedding_bags.item_emb", + ) + self.assertTrue( + np.array_equal( + view.calls[0]["data"][0]["embedding"], + np.array([1.0, 2.0], dtype=np.float32), + ) + ) + self.assertTrue( + np.array_equal( + view.calls[0]["data"][1]["embedding"], + np.array([3.0, 4.0, 5.0], dtype=np.float32), + ) + ) + self.assertEqual(view.closed, [True]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/utils/sparse_embedding_contract.py b/tzrec/utils/sparse_embedding_contract.py new file mode 100644 index 00000000..23edbb01 --- /dev/null +++ b/tzrec/utils/sparse_embedding_contract.py @@ -0,0 +1,96 @@ +# Copyright (c) 2026, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Shared sparse-embedding identity used by export, delta dump and serving.""" + +from collections import defaultdict +from dataclasses import dataclass +from typing import Dict, Iterable, Optional, Tuple + +SPARSE_EC_ROLE = "ec" +SPARSE_EBC_ROLE = "ebc" +SPARSE_EMBEDDING_ROLES = frozenset((SPARSE_EC_ROLE, SPARSE_EBC_ROLE)) +# NvEmbeddings and the future Processor consumer reserve this key as invalid. +# Other negative int64 values remain valid bit patterns for dynamic uint64 IDs. +SPARSE_EMBEDDING_INVALID_KEY = -1 + + +@dataclass(frozen=True) +class SparseEmbeddingIdentity: + """Physical sparse table identity and its cross-system canonical name.""" + + role: str + table_name: str + embedding_name: str + dimension: int + feature_names: Tuple[str, ...] = () + + +def build_sparse_embedding_name_map( + role_table_pairs: Iterable[Tuple[str, str]], +) -> Dict[Tuple[str, str], str]: + """Allocate canonical names for ``(collection role, table name)`` pairs. + + EC and EBC are separate physical collections. A table name used in only one + collection keeps its historical name. If it appears in both collections, + each physical table receives a role suffix. Candidate collisions with other + raw table names are resolved deterministically with a numeric suffix. + """ + roles_by_name = defaultdict(set) + for role, table_name in role_table_pairs: + if role not in SPARSE_EMBEDDING_ROLES: + raise ValueError(f"unsupported sparse embedding role: {role!r}") + if not table_name: + raise ValueError("sparse embedding table_name must not be empty") + roles_by_name[table_name].add(role) + + used_names = set(roles_by_name.keys()) + name_by_role_table: Dict[Tuple[str, str], str] = {} + for table_name, roles in roles_by_name.items(): + if len(roles) == 1: + role = next(iter(roles)) + name_by_role_table[(role, table_name)] = table_name + continue + + for role in sorted(roles): + base_candidate = f"{table_name}__{role}" + candidate = base_candidate + suffix = 1 + while candidate in used_names: + candidate = f"{base_candidate}_{suffix}" + suffix += 1 + used_names.add(candidate) + name_by_role_table[(role, table_name)] = candidate + return name_by_role_table + + +def resolve_sparse_embedding_name( + name_by_role_table: Dict[Tuple[str, str], str], + table_name: str, + role: Optional[str], +) -> str: + """Resolve a table to its canonical name, rejecting ambiguous identities.""" + if role is not None and (role, table_name) in name_by_role_table: + return name_by_role_table[(role, table_name)] + + candidates = [ + embedding_name + for (_role, name), embedding_name in name_by_role_table.items() + if name == table_name + ] + if len(candidates) == 1: + return candidates[0] + if not candidates: + raise KeyError(f"sparse embedding {table_name!r} is not in model metadata") + raise ValueError( + f"sparse embedding {table_name!r} appears in multiple collection kinds; " + f"cannot resolve canonical name without role, got role={role!r}" + ) diff --git a/tzrec/utils/sparse_embedding_contract_test.py b/tzrec/utils/sparse_embedding_contract_test.py new file mode 100644 index 00000000..abe1f258 --- /dev/null +++ b/tzrec/utils/sparse_embedding_contract_test.py @@ -0,0 +1,50 @@ +# Copyright (c) 2025, Alibaba Group; +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import unittest + +from tzrec.utils.sparse_embedding_contract import ( + build_sparse_embedding_name_map, + resolve_sparse_embedding_name, +) + + +class SparseEmbeddingContractTest(unittest.TestCase): + def test_single_collection_keeps_table_name(self): + names = build_sparse_embedding_name_map( + [("ec", "sequence_emb"), ("ebc", "user_emb")] + ) + self.assertEqual(names[("ec", "sequence_emb")], "sequence_emb") + self.assertEqual(names[("ebc", "user_emb")], "user_emb") + + def test_cross_collection_collision_uses_role_and_numeric_suffix(self): + names = build_sparse_embedding_name_map( + [ + ("ec", "shared"), + ("ebc", "shared"), + ("ec", "shared__ec"), + ] + ) + self.assertEqual(names[("ec", "shared")], "shared__ec_1") + self.assertEqual(names[("ebc", "shared")], "shared__ebc") + self.assertEqual(names[("ec", "shared__ec")], "shared__ec") + + def test_resolver_requires_role_for_ambiguous_table(self): + names = build_sparse_embedding_name_map([("ec", "shared"), ("ebc", "shared")]) + with self.assertRaisesRegex(ValueError, "multiple collection"): + resolve_sparse_embedding_name(names, "shared", None) + self.assertEqual( + resolve_sparse_embedding_name(names, "shared", "ec"), "shared__ec" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/version.py b/tzrec/version.py index 6939faa1..ca1bee61 100644 --- a/tzrec/version.py +++ b/tzrec/version.py @@ -9,4 +9,4 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "1.3.8" +__version__ = "1.3.9"