From 79e7f69b5f6ea5b29637f332d4edb0d2c1b8b06f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 12:34:14 +0800 Subject: [PATCH 01/21] [feat] export the prompt front-end and serving contract for genrec The HF export wrote the weights, a plain backbone config and a bare tokenizer, and refused any prompt with a projected slot, so a serving engine had to reimplement the assembler or not serve the model. It now writes the composite config a runtime resolves the backbone through, the extended tokenizer as a directory AutoTokenizer loads, prompt/prompt.json with the SID space, decode schedule and digests, and frontend/ as a standard tzrec model directory holding the scripted walk, the hole_keys fold, the trained projections and, by default, the slot tables. Under USE_DISTRIBUTED_EMBEDDING=1 the front-end becomes the processor's dense stage and the tables ship as the sparse npz files that stage already loads. The collator keeps its numpy walk; a parity test pins the two call sites of the one walk, and a multi-value SID history now assembles in both. ResolvedSidSpace gains bundle_uuid, so vocab_hash and plan_hash change once for checkpoints saved before this commit. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- docs/source/usage/export.md | 24 + tzrec/main.py | 70 +- tzrec/prompt/assembler.py | 20 +- tzrec/prompt/compile.py | 48 +- tzrec/prompt/export.py | 424 +++++++++++ tzrec/prompt/export_test.py | 249 +++++++ tzrec/prompt/frontend.py | 737 ++++++++++++++++++++ tzrec/prompt/frontend_test.py | 420 +++++++++++ tzrec/prompt/persist.py | 80 ++- tzrec/prompt/types.py | 45 +- tzrec/tests/genrec_serving_contract_test.py | 178 +++++ tzrec/tests/prompt_integration_test.py | 112 +++ tzrec/tests/prompt_test_util.py | 140 +++- tzrec/utils/export_util.py | 3 + 14 files changed, 2517 insertions(+), 33 deletions(-) create mode 100644 tzrec/prompt/export.py create mode 100644 tzrec/prompt/export_test.py create mode 100644 tzrec/prompt/frontend.py create mode 100644 tzrec/prompt/frontend_test.py create mode 100644 tzrec/tests/genrec_serving_contract_test.py diff --git a/docs/source/usage/export.md b/docs/source/usage/export.md index 734cf420..da38097f 100644 --- a/docs/source/usage/export.md +++ b/docs/source/usage/export.md @@ -221,3 +221,27 @@ eascmd -i ${ACCESS_KEY_ID} -k ${ACCESS_KEY_SECRET} -e ${ENDPOINT} create aot_exp ``` 任务运行结束后,`--export_dir` 指向的目录即为导出好的模型,将其作为在线服务的模型路径部署即可(部署方式参见 [模型服务](serving.md))。 + +(genrec-export)= + +## 生成式推荐模型(genrec)导出 + +`genrec_causal_lm_model` 导出为一个 HuggingFace 目录,供 SGLang 等 LLM 推理引擎直接加载,并附带在线 Processor 所需的 prompt 前端: + +``` +export_dir/ + config.json # 复合结构:architectures 为 PromptGenRecForCausalLM,骨干网络配置位于 text_config + model.safetensors # 骨干网络权重,参数名与骨干网络一致 + generation_config.json + prompt/ + prompt.json # 服务契约:sid_space(band、base_vocab_size、bundle_uuid)、decode 调度、vocab_hash/plan_hash、frontend 说明 + tokenizer/ # 扩展了 SID token 的 tokenizer,可由 AutoTokenizer 加载(对应 SGLang 的 --tokenizer-path) + frontend/ # 标准 tzrec 模型目录,TorchEasyRec Processor 以 model_path 直接加载 + scripted_model.pt # prompt 前端:特征 -> input_ids / hole_positions / slot_embeds / hole_keys / hole_slot_counts + fg.json pipeline.config model_acc.json +``` + +- 默认导出下,PROJECTED slot 的 embedding 表内置于 `frontend/scripted_model.pt`,前端输入为 FG 输出的原始 id(`{feature}.values` / `.lengths` / `.key_lengths`)。 +- 设置 `USE_DISTRIBUTED_EMBEDDING=1` 时,前端改为读取 Processor 分布式 embedding 阶段查表后的向量(命名由 `frontend/dense_meta.json` 描述),embedding 表以 `frontend/sparse/*.npz` 导出,格式与普通模型的分布式 embedding 导出一致。此时前端仍需要 INLINE slot 与 PROJECTED 成员的原始 id(用于计算 `hole_keys`),`prompt.json` 的 `frontend.inputs` 列出了全部输入 key。 +- 约束解码索引不由 tzrec 生成:推理侧根据 `prompt/prompt.json` 与 SID bundle 的 `sid_to_items` 构建(SGLang 侧 `python -m sglang.srt.beam_search.build_constraint_csr`),并以 `bundle_uuid` 校验索引与模型是否来自同一 bundle。 +- genrec 导出不支持 `export_config.use_dense_ema=true`。 diff --git a/tzrec/main.py b/tzrec/main.py index cba840d6..a841d8b1 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -74,7 +74,12 @@ from tzrec.optim.lr_scheduler import BaseLR from tzrec.optim.optimizer import TZRecOptimizer from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.persist import check_prompt_assets +from tzrec.prompt.persist import ( + PROMPT_DIR, + TOKENIZER_DIR, + check_prompt_assets, + write_serving_contract, +) from tzrec.prompt.types import CompiledPrompt from tzrec.protos.data_pb2 import DataConfig, DatasetType from tzrec.protos.eval_pb2 import EvalConfig @@ -1155,6 +1160,52 @@ def evaluate( logger.info("Evaluate Finished.") +def _export_prompt_serving_assets( + pipeline_config: EasyRecConfig, + features: List[BaseFeature], + compiled_prompt: CompiledPrompt, + checkpoint_path: str, + export_dir: str, +) -> None: + """Write the composite config, the front-end and the serving contract. + + The model is rebuilt and restored on CPU so the front-end carries the + trained projections and slot tables by reference, the same way the + TorchScript export restores a checkpoint before scripting it. + """ + from tzrec.prompt.export import ( + build_front_end, + export_sparse_tables, + write_composite_config, + write_front_end_dir, + ) + from tzrec.utils.state_dict_util import init_parameters + + write_composite_config(export_dir) + + model = _create_model( + pipeline_config.model_config, + features, + list(pipeline_config.data_config.label_fields), + compiled_prompt=compiled_prompt, + ) + model.set_is_inference(True) + wrapped_model = ScriptWrapper(model) + init_parameters(wrapped_model, torch.device("cpu")) + checkpoint_util.restore_model(checkpoint_path, wrapped_model) + + carry_tables = not acc_utils.use_distributed_embedding() + front_end, frontend_meta, dense_meta = build_front_end( + model, compiled_prompt, features, carry_tables + ) + frontend_dir = write_front_end_dir( + front_end, pipeline_config, features, export_dir, dense_meta + ) + if not carry_tables: + export_sparse_tables(wrapped_model, checkpoint_path, frontend_dir) + write_serving_contract(compiled_prompt, export_dir, frontend_meta) + + def export( pipeline_config_path: str, export_dir: str, @@ -1233,24 +1284,17 @@ def export( pipeline_config.prompt_config, features, list(pipeline_config.data_config.label_fields), + tokenizer_dir=os.path.join(export_dir, PROMPT_DIR, TOKENIZER_DIR) + if is_rank_zero + else None, ) check_prompt_assets(compiled_prompt, checkpoint_path) - if compiled_prompt.prompt_plan.projected_slots: - raise ValueError( - "HF export drops projected-slot state: dcp_to_hf keeps only " - "backbone keys, so embedding_group and projections would be " - "absent and the artifact could not reproduce checkpoint " - "inference." - ) if is_rank_zero: from tzrec.utils.hf_export_util import dcp_to_hf dcp_to_hf(checkpoint_path, export_dir) - compile_prompt( - pipeline_config.prompt_config, - features, - list(pipeline_config.data_config.label_fields), - tokenizer_dir=export_dir, + _export_prompt_serving_assets( + pipeline_config, features, compiled_prompt, checkpoint_path, export_dir ) return diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index 04eb64b3..b866b87e 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -279,12 +279,26 @@ def forward( ) if seg.fill is FillMode.INLINE: source, lengths = member_lengths[0] - # a dense sequence feature emits (total, value_dim); the stream - # is one code per position, so value_dim is always 1 here + lengths = np.asarray(lengths, dtype=np.int64) inline_flat[seg.name] = np.asarray( parsed_features[f"{source}.values"] ).reshape(-1) - inline_lengths[seg.name] = np.asarray(lengths, dtype=np.int64) + # a multi-value sequence holds `lengths` items per sample and + # `key_lengths` codes per item, so a sample's position count is + # a segmented sum rather than its item count + key_lengths_key = f"{source}.key_lengths" + if key_lengths_key in parsed_features: + key_lengths = np.asarray( + parsed_features[key_lengths_key], dtype=np.int64 + ).reshape(-1) + per_sample = np.zeros(lengths.size, dtype=np.int64) + np.add.at( + per_sample, + np.repeat(np.arange(lengths.size), lengths), + key_lengths, + ) + lengths = per_sample + inline_lengths[seg.name] = lengths elif seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: source, lengths = member_lengths[0] for other_source, other_lengths in member_lengths[1:]: diff --git a/tzrec/prompt/compile.py b/tzrec/prompt/compile.py index 7c54737d..030e3272 100644 --- a/tzrec/prompt/compile.py +++ b/tzrec/prompt/compile.py @@ -134,12 +134,13 @@ def _render_sid_tokens(sid_space: SidSpace) -> List[str]: return [fmt.replace("{i}", str(i)) for i in range(sum(sid_space.codebook))] -def _read_manifest_codebook(path: str) -> List[int]: - """Read ``codebook`` from a SID manifest.""" +def _read_manifest(path: str) -> Tuple[List[int], str]: + """Read ``codebook`` and the bundle identity from a SID manifest.""" if not os.path.exists(path): raise ValueError(f"sid_space.manifest_path [{path}] does not exist.") with open(path, "r") as f: - return [int(c) for c in json.load(f)["codebook"]] + manifest = json.load(f) + return [int(c) for c in manifest["codebook"]], str(manifest.get("bundle_uuid", "")) def _build_sid_space( @@ -159,8 +160,15 @@ def _build_sid_space( if any(c <= 0 for c in codebook): raise ValueError(f"every codebook size must be positive, got {codebook}.") + bundle_uuid = "" if space.HasField("manifest_path"): - declared = _read_manifest_codebook(space.manifest_path) + declared, bundle_uuid = _read_manifest(space.manifest_path) + if not bundle_uuid: + logger.warning( + f"the SID manifest at [{space.manifest_path}] carries no " + f"bundle_uuid, so serving cannot verify that a catalog belongs " + f"to this bundle by identity rather than by path." + ) if declared != codebook: raise ValueError( f"sid_space.codebook {codebook} does not match the manifest at " @@ -212,9 +220,31 @@ def _build_sid_space( sentinel_token_id=sentinel_id, eos_token_id=_special_id(tok, ("<|im_end|>", "<|endoftext|>")), pad_token_id=_special_id(tok, ("<|endoftext|>", "<|im_end|>")), + bundle_uuid=bundle_uuid, ) +def _save_tokenizer_dir( + tok: Tokenizer, sid_space: ResolvedSidSpace, tokenizer_dir: str +) -> None: + """Write the extended tokenizer as a directory ``AutoTokenizer`` loads. + + ``tokenizer.json`` carries the vocabulary and the added SID atoms; the + minimal ``tokenizer_config.json`` beside it names the tokenizer class and + the two special ids the prompt resolved, which is all a serving runtime + needs to decode a generated SID atom through ``--tokenizer-path``. + """ + os.makedirs(tokenizer_dir, exist_ok=True) + tok.save(os.path.join(tokenizer_dir, "tokenizer.json")) + config = { + "tokenizer_class": "PreTrainedTokenizerFast", + "eos_token": tok.id_to_token(sid_space.eos_token_id), + "pad_token": tok.id_to_token(sid_space.pad_token_id), + } + with open(os.path.join(tokenizer_dir, "tokenizer_config.json"), "w") as f: + json.dump(config, f, indent=2) + + def _special_id(tok: Tokenizer, candidates: Sequence[str]) -> int: """First candidate the tokenizer knows, so a family swap does not break.""" for name in candidates: @@ -247,9 +277,9 @@ def compile_prompt( cfg: the prompt config to compile. features: every feature a body slot may reference, already created. label_fields: data_config.label_fields; a response slot names these. - tokenizer_dir: where to write the extended tokenizer, flat, which is - what an exported HuggingFace directory loads; skipped when None. - Only export persists it -- nothing reads a training-time copy. + tokenizer_dir: where to write the extended tokenizer as a directory + ``AutoTokenizer`` loads; skipped when None. Only export persists + it -- nothing reads a training-time copy. Returns: The compiled prompt. @@ -330,8 +360,7 @@ def compile_prompt( sid_space = _build_sid_space(cfg, tok, base_vocab_size, has_projection) if tokenizer_dir: - os.makedirs(tokenizer_dir, exist_ok=True) - tok.save(os.path.join(tokenizer_dir, "tokenizer.json")) + _save_tokenizer_dir(tok, sid_space, tokenizer_dir) tokenizer_json = tok.to_str() slot_ids = {name: i for i, name in enumerate(resolved_slots_by_name)} @@ -374,6 +403,7 @@ def compile_prompt( logits_suffix_len=_suffix_keep(response), static_prefix_len=_static_prefix_len(body), projected_slots=projected, + slot_index={seg.name: index for index, seg in enumerate(projected)}, ) _validate(plan) diff --git a/tzrec/prompt/export.py b/tzrec/prompt/export.py new file mode 100644 index 00000000..80b266da --- /dev/null +++ b/tzrec/prompt/export.py @@ -0,0 +1,424 @@ +# 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. + +"""Writes what a serving runtime needs beside the HuggingFace weights. + +The HuggingFace config is rewritten into composite shape, with the backbone +nested under ``text_config``. That is not cosmetic: it names the backbone to a +runtime that composes an arbitrary causal LM behind one registered +architecture, and it is what makes an engine treat the model as carrying a +second modality, which is exactly what a projected slot is. + +The front-end is exported whether or not any slot is projected, as a standard +tzrec model directory so the online processor loads it unchanged: the scripted +module, ``fg.json``, ``pipeline.config`` and ``model_acc.json``. Lookup is a +stage of its own. By default the artifact carries the slot tables and takes +raw ids; under ``USE_DISTRIBUTED_EMBEDDING=1`` the processor's embedding stage +owns them, the module becomes the dense stage that reads its looked-up rows, +and the tables ship beside it as the sparse npz files that stage already +loads. +""" + +import copy +import glob +import json +import os +from typing import Any, Dict, List, Optional, Tuple, cast + +import torch +from torch import nn +from torchrec.modules.embedding_modules import ( + EmbeddingBagCollectionInterface, + EmbeddingCollectionInterface, +) +from torchrec.modules.mc_embedding_modules import ( + ManagedCollisionEmbeddingBagCollection, + ManagedCollisionEmbeddingCollection, +) + +from tzrec.acc import utils as acc_utils +from tzrec.features.feature import BaseFeature, create_feature_configs, create_fg_json +from tzrec.prompt.frontend import ( + OUT_CU_SEQLENS, + OUT_HOLE_KEYS, + OUT_HOLE_POSITIONS, + OUT_HOLE_SLOT_COUNTS, + OUT_INPUT_IDS, + OUT_RESPONSE_LENGTHS, + OUT_SLOT_EMBEDS, + PromptAssembler, + PromptFrontEnd, + SlotTable, + host_lengths_key, +) +from tzrec.prompt.types import CompiledPrompt, FillMode, SlotSeg +from tzrec.protos.model_pb2 import FeatureGroupType +from tzrec.protos.pipeline_pb2 import EasyRecConfig +from tzrec.utils import config_util +from tzrec.utils.logging_util import logger + +SERVING_ARCH = "PromptGenRecForCausalLM" +SERVING_MODEL_TYPE = "prompt_genrec" +FRONTEND_DIR = "frontend" +SCRIPTED_MODEL_FILENAME = "scripted_model.pt" +LOOKUP_ARTIFACT = "artifact" +LOOKUP_HOST = "host" +_CONFIG = "config.json" +_SPARSE_DIR = "sparse" + + +def write_composite_config(export_dir: str) -> None: + """Rewrite ``config.json`` so the backbone sits under ``text_config``. + + Args: + export_dir: the HuggingFace export directory. + """ + path = os.path.join(export_dir, _CONFIG) + with open(path, "r") as f: + backbone: Dict[str, Any] = json.load(f) + if backbone.get("model_type") == SERVING_MODEL_TYPE: + return + composite = { + "architectures": [SERVING_ARCH], + "model_type": SERVING_MODEL_TYPE, + "text_config": backbone, + } + # a runtime that reads only the outer config still needs to size its cache + for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): + if key in backbone: + composite[key] = backbone[key] + with open(path, "w") as f: + json.dump(composite, f, indent=2) + logger.info( + f"wrote a composite config naming backbone " + f"{backbone.get('architectures', ['?'])[0]} under {SERVING_ARCH}." + ) + + +def _feature_tables(embedding_group: nn.Module) -> Dict[str, Tuple[torch.Tensor, str]]: + """Map every feature with a plain table to its weight and pooling mode. + + Managed-collision tables are skipped: their ids are remapped before the + lookup, which a plain ``EmbeddingBag`` cannot reproduce. + """ + collision_prefixes = [ + name + "." + for name, module in embedding_group.named_modules() + if isinstance( + module, + ( + ManagedCollisionEmbeddingBagCollection, + ManagedCollisionEmbeddingCollection, + ), + ) + ] + tables: Dict[str, Tuple[torch.Tensor, str]] = {} + for name, module in embedding_group.named_modules(): + if any(name.startswith(prefix) for prefix in collision_prefixes): + continue + if isinstance(module, EmbeddingBagCollectionInterface): + weights = module.state_dict() + for cfg in module.embedding_bag_configs(): + weight = weights[f"embedding_bags.{cfg.name}.weight"] + for feature_name in cfg.feature_names: + tables.setdefault( + feature_name, (weight, str(cfg.pooling.value).lower()) + ) + elif isinstance(module, EmbeddingCollectionInterface): + # a multi-value item pools its codes with the feature's own + # pooling; a single-value item is that pooling over one row + weights = module.state_dict() + for cfg in module.embedding_configs(): + weight = weights[f"embeddings.{cfg.name}.weight"] + for feature_name in cfg.feature_names: + tables.setdefault(feature_name, (weight, "")) + return tables + + +def _is_multi_valued(feature: BaseFeature) -> bool: + """Whether the parser emits ``key_lengths`` for this feature.""" + return feature.is_sequence and feature.value_dim != 1 + + +def build_slot_tables( + model: nn.Module, compiled_prompt: CompiledPrompt, features: List[BaseFeature] +) -> List[SlotTable]: + """One ``SlotTable`` per projected slot, filled from the restored model. + + Args: + model: the restored genrec model, whose ``embedding_group`` holds the + trained tables. + compiled_prompt: the compiled prompt. + features: every created feature. + + Returns: + The tables in ``projected_slots`` order. + + Raises: + ValueError: a member has no plain table the artifact could carry. + """ + tables = _feature_tables(model.embedding_group) + by_name = {feature.name: feature for feature in features} + result: List[SlotTable] = [] + for seg in compiled_prompt.prompt_plan.projected_slots: + rows: List[int] = [] + dims: List[int] = [] + modes: List[str] = [] + key_length_keys: List[str] = [] + for name in seg.feature_names: + feature = by_name[name] + if not feature.is_sparse or name not in tables: + raise ValueError( + f"prompt slot [{seg.name}] member [{name}] has no plain " + f"embedding table the front-end could carry (dense, " + f"managed-collision and dynamic tables cannot be). Export " + f"with USE_DISTRIBUTED_EMBEDDING=1 so the serving host " + f"owns the lookup." + ) + weight, mode = tables[name] + rows.append(int(weight.shape[0])) + dims.append(int(weight.shape[1])) + modes.append(mode or str(feature.pooling_type.value).lower()) + key_length_keys.append( + f"{name}.key_lengths" if _is_multi_valued(feature) else "" + ) + table = SlotTable( + [f"{name}.values" for name in seg.feature_names], + [f"{name}.lengths" for name in seg.feature_names], + key_length_keys, + rows, + dims, + seg.group_type == FeatureGroupType.JAGGED_SEQUENCE, + modes, + ) + with torch.no_grad(): + for module, name in zip(table.tables, seg.feature_names): + bag = cast(nn.EmbeddingBag, module) + bag.weight.copy_(tables[name][0].detach().to(bag.weight.dtype)) + result.append(table) + return result + + +def host_embed_keys( + compiled_prompt: CompiledPrompt, +) -> Tuple[List[List[str]], Dict[str, List[str]]]: + """Batch keys a host lookup stage fills, and the ``dense_meta`` naming them. + + The names follow the dense graph a distributed-embedding export produces: + a sequence member arrives as ``{feature}`` rows with ``{feature}__lengths``, + and a pooled slot as one ``{slot}__ebc`` tensor concatenating its members. + + Returns: + Per projected slot, the keys to concatenate on the feature axis, and + the ``dense_meta.json`` the processor reads to produce them. + """ + embed_keys: List[List[str]] = [] + dense_meta: Dict[str, List[str]] = {"sequence__ec": []} + for seg in compiled_prompt.prompt_plan.projected_slots: + if seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: + embed_keys.append(list(seg.feature_names)) + for name in seg.feature_names: + dense_meta["sequence__ec"].extend( + [f"{name}__ec", host_lengths_key(name)] + ) + else: + key = f"{seg.name}__ebc" + embed_keys.append([key]) + dense_meta[key] = [f"{name}__ebc" for name in seg.feature_names] + return embed_keys, dense_meta + + +def _frontend_inputs( + compiled_prompt: CompiledPrompt, + features: List[BaseFeature], + lookup: str, + embed_keys: List[List[str]], +) -> List[str]: + """Every batch key the exported module reads, for the serving contract.""" + by_name = {feature.name: feature for feature in features} + inputs: List[str] = [] + + def raw(name: str) -> None: + inputs.append(f"{name}.values") + inputs.append(f"{name}.lengths") + if _is_multi_valued(by_name[name]): + inputs.append(f"{name}.key_lengths") + + plan = compiled_prompt.prompt_plan + for seg in plan.segments: + if isinstance(seg, SlotSeg) and seg.fill is FillMode.INLINE: + raw(seg.feature_names[0]) + for seg, keys in zip(plan.projected_slots, embed_keys): + for name in seg.feature_names: + raw(name) + if lookup == LOOKUP_HOST: + inputs.extend(keys) + if seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: + inputs.extend(host_lengths_key(name) for name in seg.feature_names) + return list(dict.fromkeys(inputs)) + + +def build_front_end( + model: nn.Module, + compiled_prompt: CompiledPrompt, + features: List[BaseFeature], + carry_tables: bool, +) -> Tuple[PromptFrontEnd, Dict[str, Any], Optional[Dict[str, List[str]]]]: + """Assemble the serving module from a restored model's parts. + + The projections are the trained modules themselves, by reference, so the + artifact carries the weights the model learned rather than a copy that can + drift from them. + + Args: + model: the restored genrec model. + compiled_prompt: the compiled prompt. + features: every created feature. + carry_tables: whether the artifact holds the slot tables, or expects a + host lookup stage to hand it rows. + + Returns: + The front-end, the ``frontend`` block of the serving contract, and the + ``dense_meta`` a host lookup stage needs (None when tables are carried). + """ + plan = compiled_prompt.prompt_plan + sid_space = compiled_prompt.sid_space + assembler = PromptAssembler( + plan, + sid_space, + features_are_dense={f.name: not f.is_sparse for f in features}, + features_are_multi_valued={f.name: _is_multi_valued(f) for f in features}, + plan_hash=compiled_prompt.plan_hash, + # serving has no answer to assemble; the LM generates it + include_response=False, + ) + projections = list(model._slot_projections) + if carry_tables: + lookup = LOOKUP_ARTIFACT + tables: Optional[List[SlotTable]] = build_slot_tables( + model, compiled_prompt, features + ) + embed_keys: List[List[str]] = [[] for _ in plan.projected_slots] + dense_meta: Optional[Dict[str, List[str]]] = None + else: + lookup = LOOKUP_HOST + tables = None + embed_keys, dense_meta = host_embed_keys(compiled_prompt) + front_end = PromptFrontEnd( + assembler, + projections, + embed_keys, + tables=tables, + vocab_hash=compiled_prompt.vocab_hash, + plan_hash=compiled_prompt.plan_hash, + bundle_uuid=sid_space.bundle_uuid if sid_space is not None else "", + ) + meta = { + "dir": FRONTEND_DIR, + "model": SCRIPTED_MODEL_FILENAME, + "lookup": lookup, + "inputs": _frontend_inputs(compiled_prompt, features, lookup, embed_keys), + "outputs": [ + OUT_INPUT_IDS, + OUT_CU_SEQLENS, + OUT_HOLE_POSITIONS, + OUT_HOLE_KEYS, + OUT_HOLE_SLOT_COUNTS, + OUT_SLOT_EMBEDS, + OUT_RESPONSE_LENGTHS, + ], + } + return front_end, meta, dense_meta + + +def write_front_end_dir( + front_end: PromptFrontEnd, + pipeline_config: EasyRecConfig, + features: List[BaseFeature], + export_dir: str, + dense_meta: Optional[Dict[str, List[str]]] = None, +) -> str: + """Script the front-end into ``frontend/``, laid out as a tzrec model dir. + + Args: + front_end: the module to export. + pipeline_config: the pipeline config, whose feature configs are + rewritten beside the copied assets. + features: every created feature. + export_dir: the HuggingFace export directory. + dense_meta: the processor's dense-stage input map, written when a host + lookup stage feeds the module. + + Returns: + The ``frontend/`` directory written. + """ + frontend_dir = os.path.join(export_dir, FRONTEND_DIR) + os.makedirs(frontend_dir, exist_ok=True) + torch.jit.script(front_end.eval()).save( + os.path.join(frontend_dir, SCRIPTED_MODEL_FILENAME) + ) + + feature_configs = create_feature_configs(features, asset_dir=frontend_dir) + served_config = copy.deepcopy(pipeline_config) + served_config.ClearField("feature_configs") + served_config.feature_configs.extend(feature_configs) + config_util.save_message( + served_config, os.path.join(frontend_dir, "pipeline.config") + ) + with open(os.path.join(frontend_dir, "fg.json"), "w") as f: + json.dump(create_fg_json(features, asset_dir=frontend_dir), f, indent=4) + with open(os.path.join(frontend_dir, "model_acc.json"), "w") as f: + json.dump(acc_utils.export_acc_config(), f, indent=4) + if dense_meta is not None: + with open(os.path.join(frontend_dir, "dense_meta.json"), "w") as f: + json.dump(dense_meta, f, indent=4) + logger.info(f"wrote the scripted prompt front-end to {frontend_dir}.") + return frontend_dir + + +def export_sparse_tables( + wrapped_model: nn.Module, checkpoint_path: str, frontend_dir: str +) -> None: + """Write the slot tables as the sparse npz files a host lookup stage loads. + + Single rank, in the layout ``export_distributed_embedding`` produces. + + Args: + wrapped_model: the restored model under its inference wrapper, so the + table names match the checkpoint's. + checkpoint_path: the checkpoint, read for dynamic tables. + frontend_dir: the ``frontend/`` directory to write ``sparse/`` into. + """ + from tzrec.utils import export_util, npz_util + + bag_info, emb_info = export_util._get_sparse_table_to_embedding_info(wrapped_model) + local, dynamic, emb_meta, feat_meta = export_util._get_sparse_embedding_tensor( + wrapped_model, checkpoint_path, emb_info, bag_info + ) + sparse_dir = os.path.join(frontend_dir, _SPARSE_DIR) + os.makedirs(sparse_dir, exist_ok=True) + shard = "sparse_embeddings-00-of-01" + npz_util.savez_streaming(os.path.join(sparse_dir, f"{shard}.npz"), local) + if dynamic: + npz_util.savez_streaming( + os.path.join(sparse_dir, "sparse_dynamic_embedding-00-of-01.npz"), dynamic + ) + with open(os.path.join(sparse_dir, f"{shard}.json"), "w") as f: + json.dump(emb_meta, f, indent=4) + with open(os.path.join(sparse_dir, "sparse_features.json"), "w") as f: + json.dump(feat_meta, f, indent=4) + shards = [] + for path in glob.glob(os.path.join(sparse_dir, "sparse_embeddings*.json")): + with open(path, "r") as f: + shards.append(json.load(f)) + with open(os.path.join(sparse_dir, "sparse_embedding.json"), "w") as f: + json.dump(export_util._merge_sharded_embedding_json(shards), f, indent=4) + logger.info(f"wrote {len(local)} slot table(s) to {sparse_dir}.") diff --git a/tzrec/prompt/export_test.py b/tzrec/prompt/export_test.py new file mode 100644 index 00000000..204caaef --- /dev/null +++ b/tzrec/prompt/export_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 json +import os +import unittest + +import torch +from torchrec import KeyedJaggedTensor + +from tzrec.datasets.utils import BASE_DATA_GROUP, Batch +from tzrec.prompt.export import ( + LOOKUP_ARTIFACT, + LOOKUP_HOST, + build_front_end, + build_slot_tables, + host_embed_keys, + write_composite_config, + write_front_end_dir, +) +from tzrec.prompt.persist import write_serving_contract +from tzrec.protos.pipeline_pb2 import EasyRecConfig +from tzrec.tests.prompt_test_util import ( + _HIST, + GenrecModelTestBase, + create_prompt_feature, + offset_sid_codes, + projected_feature, +) +from tzrec.utils.state_dict_util import init_parameters + +_CAT = ( + 'id_feature { feature_name: "cat" expression: "user:cat" num_buckets: 16 ' + 'embedding_dim: 8 pooling: "mean" }' +) +_VEC = 'raw_feature { feature_name: "vec" expression: "user:vec" value_dim: 4 }' +_TEMPLATE = "History : {{hist}} . {{beh}} {{cat}} Predict :" + + +class PromptExportTest(GenrecModelTestBase): + """Builds the serving pieces from an in-memory model with projected slots.""" + + def setUp(self) -> None: + super().setUp() + self.features = [ + create_prompt_feature(_HIST), + create_prompt_feature(projected_feature("beh", 8)), + create_prompt_feature(_CAT), + ] + self.compiled_prompt = self._compile(self.features, template=_TEMPLATE) + self.model = self._model() + init_parameters(self.model, torch.device("cpu")) + self.raw = { + "hist.values": torch.tensor( + offset_sid_codes([0, 1, 2, 3, 0, 1], [4, 4, 4]) + ), + "hist.lengths": torch.tensor([6]), + "beh.values": torch.tensor([3, 9]), + "beh.lengths": torch.tensor([2]), + "cat.values": torch.tensor([5]), + "cat.lengths": torch.tensor([1]), + } + sparse = KeyedJaggedTensor( + keys=["beh", "cat"], + values=torch.tensor([3, 9, 5]), + lengths=torch.tensor([2, 1]), + ) + self.grouped = self.model.embedding_group( + Batch(sparse_features={BASE_DATA_GROUP: sparse}) + ) + + def _expected_slot_embeds(self) -> torch.Tensor: + projections = self.model._slot_projections + return torch.cat( + [ + projections[0](self.grouped["beh.sequence"]), + projections[1](self.grouped["cat"]), + ] + ) + + def test_slot_tables_reproduce_the_embedding_group(self) -> None: + """The carried tables look up exactly what the training model did.""" + tables = build_slot_tables(self.model, self.compiled_prompt, self.features) + self.assertEqual(len(tables), 2) + self.assertTrue( + torch.allclose(tables[0](self.raw), self.grouped["beh.sequence"]) + ) + self.assertTrue(torch.allclose(tables[1](self.raw), self.grouped["cat"])) + + def test_the_artifact_shape_matches_the_model_on_the_same_batch(self) -> None: + front_end, meta, dense_meta = build_front_end( + self.model, self.compiled_prompt, self.features, carry_tables=True + ) + out = front_end(self.raw) + self.assertTrue( + torch.allclose(out["slot_embeds"], self._expected_slot_embeds(), atol=1e-6) + ) + self.assertEqual(out["hole_slot_counts"].tolist(), [2, 1]) + self.assertEqual(meta["lookup"], LOOKUP_ARTIFACT) + self.assertIsNone(dense_meta) + self.assertEqual( + meta["inputs"], + [ + "hist.values", + "hist.lengths", + "beh.values", + "beh.lengths", + "cat.values", + "cat.lengths", + ], + ) + self.assertEqual(front_end.vocab_hash, self.compiled_prompt.vocab_hash) + + def test_the_host_shape_reads_the_processor_keys(self) -> None: + """A host lookup stage hands over rows named as dense_meta describes.""" + artifact, _, _ = build_front_end( + self.model, self.compiled_prompt, self.features, carry_tables=True + ) + host, meta, dense_meta = build_front_end( + self.model, self.compiled_prompt, self.features, carry_tables=False + ) + self.assertEqual(meta["lookup"], LOOKUP_HOST) + self.assertEqual( + dense_meta, + {"sequence__ec": ["beh__ec", "beh__lengths"], "cat__ebc": ["cat__ebc"]}, + ) + self.assertEqual( + host_embed_keys(self.compiled_prompt)[0], [["beh"], ["cat__ebc"]] + ) + batch = dict(self.raw) + batch["beh"] = self.grouped["beh.sequence"] + batch["beh__lengths"] = self.raw["beh.lengths"] + batch["cat__ebc"] = self.grouped["cat"] + expected = artifact(self.raw) + got = torch.jit.script(host.eval())(batch) + self.assertTrue(torch.allclose(got["slot_embeds"], expected["slot_embeds"])) + self.assertTrue(torch.equal(got["hole_keys"], expected["hole_keys"])) + self.assertIn("beh__lengths", meta["inputs"]) + self.assertIn("cat__ebc", meta["inputs"]) + + def test_a_member_without_a_plain_table_cannot_be_carried(self) -> None: + features = [create_prompt_feature(_HIST), create_prompt_feature(_VEC)] + compiled = self._compile(features, template="History : {{hist}} {{vec}} :") + model = self._model(features=features, compiled_prompt=compiled) + init_parameters(model, torch.device("cpu")) + with self.assertRaisesRegex(ValueError, "USE_DISTRIBUTED_EMBEDDING=1"): + build_slot_tables(model, compiled, features) + + def test_write_composite_config_is_idempotent(self) -> None: + path = os.path.join(self.test_dir, "config.json") + backbone = { + "model_type": "qwen2", + "architectures": ["Qwen2ForCausalLM"], + "hidden_size": 32, + "vocab_size": 77, + } + with open(path, "w") as f: + json.dump(backbone, f) + write_composite_config(self.test_dir) + write_composite_config(self.test_dir) + with open(path, "r") as f: + composite = json.load(f) + self.assertEqual(composite["architectures"], ["PromptGenRecForCausalLM"]) + self.assertEqual(composite["model_type"], "prompt_genrec") + self.assertEqual(composite["text_config"], backbone) + self.assertEqual(composite["hidden_size"], 32) + + def test_write_front_end_dir_lays_out_a_tzrec_model_dir(self) -> None: + front_end, _, _ = build_front_end( + self.model, self.compiled_prompt, self.features, carry_tables=True + ) + frontend_dir = write_front_end_dir( + front_end, + EasyRecConfig(model_dir="unused"), + self.features, + self.test_dir, + dense_meta={"sequence__ec": ["beh__ec", "beh__lengths"]}, + ) + for name in ( + "scripted_model.pt", + "fg.json", + "pipeline.config", + "model_acc.json", + "dense_meta.json", + ): + self.assertTrue(os.path.exists(os.path.join(frontend_dir, name)), name) + with open(os.path.join(frontend_dir, "model_acc.json"), "r") as f: + self.assertEqual(json.load(f)["SPARSE_INT64"], "1") + with open(os.path.join(frontend_dir, "fg.json"), "r") as f: + names = [feature["feature_name"] for feature in json.load(f)["features"]] + self.assertEqual(names, ["hist", "beh", "cat"]) + loaded = torch.jit.load(os.path.join(frontend_dir, "scripted_model.pt")) + out = loaded(self.raw, torch.device("cpu")) + self.assertTrue( + torch.allclose(out["slot_embeds"], self._expected_slot_embeds(), atol=1e-6) + ) + + def test_serving_contract_lists_what_the_runtime_reads(self) -> None: + _, meta, _ = build_front_end( + self.model, self.compiled_prompt, self.features, carry_tables=True + ) + path = write_serving_contract(self.compiled_prompt, self.test_dir, meta) + with open(path, "r") as f: + contract = json.load(f) + space = self.compiled_prompt.sid_space + self.assertEqual(contract["sid_space"]["band_lo"], list(space.band_lo)) + self.assertEqual( + contract["sid_space"]["base_vocab_size"], space.base_vocab_size + ) + self.assertEqual(contract["sid_space"]["bundle_uuid"], "") + self.assertEqual( + contract["decode"], + [ + { + "level": level, + "band_lo": space.band_lo[level], + "band_hi": space.band_hi[level], + } + for level in range(3) + ], + ) + self.assertEqual(contract["vocab_hash"], self.compiled_prompt.vocab_hash) + self.assertEqual(contract["plan_hash"], self.compiled_prompt.plan_hash) + self.assertEqual(contract["frontend"], meta) + self.assertEqual(contract["static_prefix_len"], 2) + + def test_a_static_run_in_the_response_is_a_forced_token(self) -> None: + compiled = self._compile( + self.features, template=_TEMPLATE, response="Predict : {{answer}}" + ) + path = write_serving_contract(compiled, self.test_dir, {}) + with open(path, "r") as f: + decode = json.load(f)["decode"] + self.assertEqual(len(decode), 5) + self.assertEqual(decode[0], {"token_id": 1}) + self.assertEqual(decode[1], {"token_id": 2}) + self.assertIn("band_lo", decode[2]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/prompt/frontend.py b/tzrec/prompt/frontend.py new file mode 100644 index 00000000..338d476b --- /dev/null +++ b/tzrec/prompt/frontend.py @@ -0,0 +1,737 @@ +# 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. + +"""The prompt front-end a serving runtime loads: assemble, look up, project. + +The assembler is the collator's walk as tensor ops. ``PromptPlan`` is a +compile-time constant, so the segment loop unrolls into parallel constant +lists at construction and what remains is jagged integer arithmetic -- +``cumsum``, ``repeat_interleave``, ``index_copy_`` -- with no data-dependent +control flow, which is what lets ``torch.jit.script`` carry it into a runtime +that has no tzrec source. + +``hole_keys`` is integer end to end: ``int64`` addition is associative and +commutative and wraps deterministically, so the fold cannot depend on the order +``index_add_`` happens to reduce in, on the device, or on how the batch was +split. A float accumulator would satisfy "do not fold the projected vector" in +letter and reintroduce the variance in spirit. +""" + +from typing import Dict, Final, List, Optional, Tuple + +import torch +from torch import nn + +from tzrec.prompt.types import ( + FillMode, + FoldConstants, + PromptPlan, + ResolvedSidSpace, + SlotSeg, + Static, +) +from tzrec.protos.model_pb2 import FeatureGroupType + +OUT_INPUT_IDS = "input_ids" +OUT_CU_SEQLENS = "cu_seqlens" +OUT_HOLE_POSITIONS = "hole_positions" +OUT_HOLE_KEYS = "hole_keys" +OUT_HOLE_SLOT_COUNTS = "hole_slot_counts" +OUT_SLOT_EMBEDS = "slot_embeds" +OUT_RESPONSE_LENGTHS = "response_lengths" + + +@torch.jit.script +def mix64(z: torch.Tensor) -> torch.Tensor: + """A SplitMix64-shaped avalanche over int64, wrapping. + + torch's right shift on a signed integer is arithmetic, so every shift the + mixer wants as logical is masked back. Getting that wrong is not a weaker + hash, it is a different function on negative inputs. + """ + z = z * (-7046029254386353131) + z = (z ^ ((z >> 30) & 0x3FFFFFFFF)) * (-4658895280553007687) + z = (z ^ ((z >> 27) & 0x1FFFFFFFFF)) * (-7723592293110705685) + return z ^ ((z >> 31) & 0x1FFFFFFFF) + + +@torch.jit.script +def _exclusive_cumsum(values: torch.Tensor) -> torch.Tensor: + """Exclusive prefix sum along dim 0.""" + return torch.cumsum(values, dim=0) - values + + +@torch.jit.script +def _row_ids(lengths: torch.Tensor) -> torch.Tensor: + """Row index of every element in a jagged buffer described by lengths.""" + return torch.repeat_interleave( + torch.arange(lengths.numel(), dtype=torch.int64, device=lengths.device), + lengths, + ) + + +@torch.jit.script +def _within_row_index(lengths: torch.Tensor) -> torch.Tensor: + """Position of every element inside its own row.""" + total = int(torch.sum(lengths)) + starts = _exclusive_cumsum(lengths) + return torch.arange( + total, dtype=torch.int64, device=lengths.device + ) - torch.repeat_interleave(starts, lengths) + + +@torch.jit.script +def _pick_device(batch: Dict[str, torch.Tensor]) -> torch.device: + """Device the batch already lives on, so nothing is built on the wrong one.""" + for value in batch.values(): + return value.device + return torch.device("cpu") + + +@torch.jit.script +def _destinations(seg_start: torch.Tensor, seg_len: torch.Tensor) -> torch.Tensor: + """Absolute index of every value of one segment in the packed stream.""" + return torch.repeat_interleave(seg_start, seg_len) + _within_row_index(seg_len) + + +def _wrap64(value: int) -> int: + """Reduce a Python int into the signed 64-bit range, wrapping. + + Every constant the fold mixes is built on the host and handed to torch as a + scalar, and torch refuses a scalar outside the tensor's dtype. Wrapping here + is what makes "the fold wraps" true at the boundary as well as inside it. + """ + value &= (1 << 64) - 1 + return value - (1 << 64) if value >= 1 << 63 else value + + +def _plan_salt(fold: FoldConstants, plan_hash: str) -> int: + """Low 64 bits of ``plan_hash``, as a signed multiple of the plan constant.""" + if not plan_hash: + return 0 + return _wrap64(_wrap64(int(plan_hash[:16], 16)) * fold.plan) + + +def host_lengths_key(feature_name: str) -> str: + """The per-row item count a host lookup stage emits beside its rows.""" + return f"{feature_name}__lengths" + + +class PromptAssembler(nn.Module): + """Walks a compiled plan to build one batch's packed token stream. + + Args: + plan: the compiled walk order and its constants. + sid_space: resolved SID token space; required when a slot renders SID + codes or is projected. + features_are_dense: feature name to whether it arrives as floats. + features_are_multi_valued: feature name to whether it carries + ``key_lengths``, that is whether ``value_dim != 1``. + plan_hash: the compiled plan's hash; its low bits salt every key, so + two plans cannot cross-match in a shared prefix cache. + include_response: whether to emit the supervised tail. + """ + + # TorchScript resolves a Final class attribute as a constant; a + # module-level one it cannot see at all + KIND_STATIC: Final[int] = 0 + KIND_INLINE: Final[int] = 1 + KIND_PROJECTED: Final[int] = 2 + # member index and position within a hole packed into one integer, wide + # enough that a position cannot carry into the member index + MEMBER_STRIDE: Final[int] = 1 << 32 + + kinds: List[int] + static_tokens: List[List[int]] + value_keys: List[str] + length_keys: List[str] + host_length_keys: List[str] + key_length_keys: List[str] + hole_slots: List[int] + slot_ids: List[int] + is_sequences: List[bool] + salts: List[int] + member_value_keys: List[List[str]] + member_length_keys: List[List[str]] + member_host_length_keys: List[List[str]] + member_key_length_keys: List[List[str]] + member_is_dense: List[bool] + + def __init__( + self, + plan: PromptPlan, + sid_space: Optional[ResolvedSidSpace] = None, + features_are_dense: Optional[Dict[str, bool]] = None, + features_are_multi_valued: Optional[Dict[str, bool]] = None, + plan_hash: str = "", + include_response: bool = True, + ) -> None: + super().__init__() + dense = features_are_dense if features_are_dense is not None else {} + multi = ( + features_are_multi_valued if features_are_multi_valued is not None else {} + ) + segments = tuple(plan.segments) + self.num_body = len(segments) + if include_response: + segments = segments + tuple(plan.response_segments) + + sentinel = -1 + if sid_space is not None and sid_space.sentinel_token_id is not None: + sentinel = int(sid_space.sentinel_token_id) + self.sentinel = sentinel + id_shift = 0 if sid_space is None else int(sid_space.base_vocab_size) + + self.kinds = [] + self.static_tokens = [] + self.value_keys = [] + self.length_keys = [] + self.host_length_keys = [] + self.key_length_keys = [] + self.hole_slots = [] + self.slot_ids = [] + self.is_sequences = [] + self.member_value_keys = [] + self.member_length_keys = [] + self.member_host_length_keys = [] + self.member_key_length_keys = [] + self.member_is_dense = [] + + for seg in segments: + if isinstance(seg, Static): + self._append( + PromptAssembler.KIND_STATIC, + [int(t) for t in seg.token_ids], + "", + "", + "", + "", + -1, + -1, + False, + [], + [], + [], + [], + ) + continue + + assert isinstance(seg, SlotSeg) + primary = seg.feature_names[0] + is_sequence = seg.group_type == FeatureGroupType.JAGGED_SEQUENCE + if seg.fill is FillMode.INLINE: + if sid_space is None: + raise ValueError( + f"prompt slot [{seg.name}] renders INLINE, which means " + f"SID codes, but no sid_space was compiled." + ) + self._append( + PromptAssembler.KIND_INLINE, + [], + f"{primary}.values", + f"{primary}.lengths", + host_lengths_key(primary), + f"{primary}.key_lengths" if multi.get(primary, False) else "", + -1, + int(seg.slot_id), + is_sequence, + [], + [], + [], + [], + ) + continue + + if sentinel < 0: + raise ValueError( + f"prompt slot [{seg.name}] is PROJECTED but no sentinel token " + f"was compiled; a hole would be indistinguishable from content." + ) + for name in seg.feature_names: + if dense.get(name, False) and is_sequence: + raise ValueError( + f"prompt slot [{seg.name}] member [{name}] is a dense " + f"sequence; the fold has no per-item boundary for it." + ) + self._append( + PromptAssembler.KIND_PROJECTED, + [], + "", + f"{primary}.lengths", + host_lengths_key(primary), + "", + int(plan.slot_index[seg.name]), + int(seg.slot_id), + is_sequence, + [f"{name}.values" for name in seg.feature_names], + [f"{name}.lengths" for name in seg.feature_names], + [host_lengths_key(name) for name in seg.feature_names], + [ + f"{name}.key_lengths" if multi.get(name, False) else "" + for name in seg.feature_names + ], + ) + self.member_is_dense = self.member_is_dense + [ + dense.get(name, False) for name in seg.feature_names + ] + + self.id_shift = id_shift + self.num_segments = len(self.kinds) + self.num_hole_slots = len(plan.projected_slots) + self.fold_value = int(plan.fold.value) + self.fold_index = int(plan.fold.index) + plan_salt = _plan_salt(plan.fold, plan_hash) + self.salts = [ + _wrap64(plan.fold.slot * slot_id + plan_salt) for slot_id in self.slot_ids + ] + + def _append( + self, + kind: int, + tokens: List[int], + value_key: str, + length_key: str, + host_length_key: str, + key_length_key: str, + hole_slot: int, + slot_id: int, + is_sequence: bool, + member_values: List[str], + member_lengths: List[str], + member_host_lengths: List[str], + member_key_lengths: List[str], + ) -> None: + """Record one unrolled segment's constants.""" + self.kinds.append(kind) + self.static_tokens.append(tokens) + self.value_keys.append(value_key) + self.length_keys.append(length_key) + self.host_length_keys.append(host_length_key) + self.key_length_keys.append(key_length_key) + self.hole_slots.append(hole_slot) + self.slot_ids.append(slot_id) + self.is_sequences.append(is_sequence) + self.member_value_keys.append(member_values) + self.member_length_keys.append(member_lengths) + self.member_host_length_keys.append(member_host_lengths) + self.member_key_length_keys.append(member_key_lengths) + + def _lengths( + self, batch: Dict[str, torch.Tensor], key: str, host_key: str + ) -> torch.Tensor: + """Per-row item count, from the parsed dict or from a host lookup stage.""" + if key in batch: + return batch[key].to(torch.int64) + return batch[host_key].to(torch.int64) + + def _batch_size(self, batch: Dict[str, torch.Tensor]) -> int: + """Row count, from the first segment that names a per-row length.""" + for i in range(self.num_segments): + key = self.length_keys[i] + if key != "": + if key in batch: + return int(batch[key].numel()) + host_key = self.host_length_keys[i] + if host_key in batch: + return int(batch[host_key].numel()) + if "batch_size" in batch: + return int(batch["batch_size"]) + return 0 + + def _values_per_row( + self, batch: Dict[str, torch.Tensor], index: int, batch_size: int + ) -> torch.Tensor: + """Token count each row contributes for one INLINE segment. + + For a multi-value sequence -- which is what a SID history is -- the row + holds ``lengths`` items and each item holds ``key_lengths`` values, so + the count is a segmented sum rather than ``lengths`` itself. + """ + lengths = self._lengths( + batch, self.length_keys[index], self.host_length_keys[index] + ) + key = self.key_length_keys[index] + if key == "": + return lengths + key_lengths = batch[key].to(torch.int64) + counts = torch.zeros(batch_size, dtype=torch.int64, device=lengths.device) + return counts.index_add_(0, _row_ids(lengths), key_lengths) + + def _segment( + self, + batch: Dict[str, torch.Tensor], + index: int, + batch_size: int, + device: torch.device, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-row length and row-major values of one segment.""" + kind = self.kinds[index] + + if kind == self.KIND_STATIC: + run = torch.tensor( + self.static_tokens[index], dtype=torch.int64, device=device + ) + width = run.numel() + seg_len = torch.full((batch_size,), width, dtype=torch.int64, device=device) + return seg_len, run.unsqueeze(0).expand(batch_size, width).reshape(-1) + + if kind == self.KIND_INLINE: + seg_len = self._values_per_row(batch, index, batch_size) + values = batch[self.value_keys[index]].to(torch.int64).reshape(-1) + return seg_len, values + self.id_shift + + if self.is_sequences[index]: + seg_len = self._lengths( + batch, self.length_keys[index], self.host_length_keys[index] + ) + else: + seg_len = torch.ones(batch_size, dtype=torch.int64, device=device) + total = int(torch.sum(seg_len)) + return seg_len, torch.full( + (total,), self.sentinel, dtype=torch.int64, device=seg_len.device + ) + + def _fold_segment( + self, + batch: Dict[str, torch.Tensor], + index: int, + batch_size: int, + hole_base: int, + keys: torch.Tensor, + ) -> None: + """Mix one projected segment's input values into ``keys``. + + Every member value that produces a hole contributes, discriminated by + slot, by member and by its index inside the hole. Without the last two + a two-member slot with values ``(a, b)`` would match one with + ``(b, a)``, and a permuted multi-value item would match itself + reordered -- both plausible, both wrong, and both silent. + """ + salt = self.salts[index] + member_values = self.member_value_keys[index] + for member in range(len(member_values)): + raw = batch[member_values[member]] + lengths = self._lengths( + batch, + self.member_length_keys[index][member], + self.member_host_length_keys[index][member], + ) + key_length_key = self.member_key_length_keys[index][member] + + if not self.is_sequences[index]: + # one hole per row; a dense member contributes its float32 bit + # pattern verbatim, which is the parsed input and not a + # computed reduction, so it is stable for a given request + if raw.dtype == torch.float32: + width = raw.size(1) + values = ( + raw.contiguous().view(torch.int32).to(torch.int64).reshape(-1) + & 0xFFFFFFFF + ) + hole = torch.repeat_interleave( + torch.arange(batch_size, dtype=torch.int64, device=raw.device), + torch.full( + (batch_size,), width, dtype=torch.int64, device=raw.device + ), + ) + local = ( + torch.arange(width, dtype=torch.int64, device=raw.device) + .unsqueeze(0) + .expand(batch_size, width) + .reshape(-1) + ) + else: + values = raw.to(torch.int64).reshape(-1) + hole = _row_ids(lengths) + local = _within_row_index(lengths) + else: + values = raw.to(torch.int64).reshape(-1) + if key_length_key == "": + hole = torch.arange( + values.numel(), dtype=torch.int64, device=values.device + ) + local = torch.zeros_like(hole) + else: + key_lengths = batch[key_length_key].to(torch.int64) + hole = _row_ids(key_lengths) + local = _within_row_index(key_lengths) + + local = local + member * self.MEMBER_STRIDE + mixed = mix64(values * self.fold_value + local * self.fold_index + salt) + keys.index_add_(0, hole + hole_base, mixed) + + def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + """Assemble one batch. + + Args: + batch: the parsed feature dict, keyed ``{feature}.values`` / + ``.lengths`` / ``.key_lengths`` as both hosts already emit it. + + Returns: + ``input_ids``, ``cu_seqlens``, ``hole_positions``, ``hole_keys``, + ``hole_slot_counts`` and ``response_lengths``. The two hole-indexed + outputs are row-aligned: entry ``k`` of each describes the same + hole, in ``projected_slots`` order. + """ + batch_size = self._batch_size(batch) + device = _pick_device(batch) + + seg_lens: List[torch.Tensor] = [] + seg_values: List[torch.Tensor] = [] + for i in range(self.num_segments): + length, values = self._segment(batch, i, batch_size, device) + seg_lens.append(length) + seg_values.append(values) + + stacked = torch.stack(seg_lens, dim=0) + row_total = torch.sum(stacked, dim=0) + row_start = _exclusive_cumsum(row_total) + seg_offsets = torch.cumsum(stacked, dim=0) - stacked + + total_tokens = int(torch.sum(row_total)) + out = torch.zeros(total_tokens, dtype=torch.int64, device=row_total.device) + + # one entry per projected slot, and an empty stream under Pattern I, + # where the concatenation below would otherwise have nothing to join + hole_parts: List[torch.Tensor] = [ + torch.zeros(0, dtype=torch.int64, device=device) + ] + for _ in range(self.num_hole_slots): + hole_parts.append(torch.zeros(0, dtype=torch.int64, device=device)) + + response_lengths = torch.zeros( + batch_size, dtype=torch.int64, device=row_total.device + ) + for i in range(self.num_segments): + dest = _destinations(row_start + seg_offsets[i], seg_lens[i]) + out.index_copy_(0, dest, seg_values[i]) + slot = self.hole_slots[i] + if slot >= 0: + hole_parts[slot + 1] = dest + if i >= self.num_body: + response_lengths = response_lengths + seg_lens[i] + + hole_positions = torch.cat(hole_parts, dim=0) + # how many of those holes each projected slot owns, so a host can cut + # the flat streams back into per-slot spans without re-deriving the plan + slot_counts = torch.zeros(self.num_hole_slots, dtype=torch.int64, device=device) + for slot in range(self.num_hole_slots): + slot_counts[slot] = hole_parts[slot + 1].numel() + keys = torch.zeros(hole_positions.numel(), dtype=torch.int64, device=out.device) + hole_base = 0 + for slot in range(self.num_hole_slots): + for i in range(self.num_segments): + if self.hole_slots[i] == slot: + self._fold_segment(batch, i, batch_size, hole_base, keys) + hole_base = hole_base + hole_parts[slot + 1].numel() + + cu_seqlens = torch.cat( + [ + torch.zeros(1, dtype=torch.int64, device=row_total.device), + torch.cumsum(row_total, dim=0), + ] + ) + # literals rather than the module constants above: TorchScript cannot + # see a module-level global. ``frontend_test`` pins the two together. + return { + "input_ids": out, + "cu_seqlens": cu_seqlens.to(torch.int32), + "hole_positions": hole_positions, + "hole_keys": keys, + "hole_slot_counts": slot_counts, + "response_lengths": response_lengths, + } + + +class SlotTable(nn.Module): + """One projected slot's members, as plain embedding tables. + + Used when the host does not run an embedding stage of its own. The lookup + has to live somewhere: either the host does it and hands over vectors, or + the artifact carries the tables. Keeping both shapes behind one module means + the walk, the fold and the projections do not change between them. + + Args: + value_keys: each member's value key. + length_keys: each member's per-row counts. + key_length_keys: each member's per-item value counts, empty when the + member holds one value per item. + num_embeddings: each member's table size. + dims: each member's embedding dimension. + is_sequence: whether holes are items rather than rows. + modes: each member's pooling, ``sum`` or ``mean``. + """ + + value_keys: List[str] + length_keys: List[str] + key_length_keys: List[str] + + def __init__( + self, + value_keys: List[str], + length_keys: List[str], + key_length_keys: List[str], + num_embeddings: List[int], + dims: List[int], + is_sequence: bool, + modes: Optional[List[str]] = None, + ) -> None: + super().__init__() + self.value_keys = value_keys + self.length_keys = length_keys + self.key_length_keys = key_length_keys + self.is_sequence = is_sequence + modes = modes if modes is not None else ["sum"] * len(dims) + self.tables = nn.ModuleList( + [ + nn.EmbeddingBag(rows, dim, mode=mode, include_last_offset=True) + for rows, dim, mode in zip(num_embeddings, dims, modes) + ] + ) + + def forward(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor: + """Look every member up and concatenate along the feature axis. + + Returns: + ``(num_holes, group_total_dim)``, ordered by hole. + """ + parts: List[torch.Tensor] = [] + index = 0 + for table in self.tables: + values = batch[self.value_keys[index]].to(torch.int64).reshape(-1) + key = self.key_length_keys[index] + if key != "": + counts = batch[key].to(torch.int64) + elif self.is_sequence: + counts = torch.ones_like(values) + else: + counts = batch[self.length_keys[index]].to(torch.int64) + offsets = torch.cat( + [ + torch.zeros(1, dtype=torch.int64, device=values.device), + torch.cumsum(counts, dim=0), + ] + ) + parts.append(table(values, offsets)) + index += 1 + return torch.cat(parts, dim=1) + + +class PromptFrontEnd(nn.Module): + """The serving artifact: assemble, look up, project. + + Exported whether or not any slot is projected -- with none it degenerates + to the assembler and three empty streams, which is cheap and keeps what a + prompt happens to contain from deciding whether a serving-critical artifact + exists. + + ``slot_embeds`` is ordered by ascending hole, matching ``hole_positions`` + entry for entry. That ordering is an obligation rather than a check: the + engine scatters positionally, so a permuted source gives every hole a + neighbour's embedding, the counts still agree, and nothing raises. + + Lookup is a stage of its own. Without a host embedding stage the artifact + carries the tables and the batch holds raw ids. Under a host stage the + embeddings are already in the batch when this runs, one entry per key in + ``embed_keys`` for each slot, concatenated on the feature axis. Everything + downstream of the lookup is identical either way. + + Args: + assembler: the walk. + projections: one per projected slot in ``projected_slots`` order; slots + sharing a module appear more than once, by reference. + embed_keys: per projected slot, the batch keys holding its members' + looked-up rows, in member order. Ignored when ``tables`` is given. + tables: one per projected slot, when the artifact carries the slot + tables rather than receiving vectors from a host stage. + vocab_hash: the compiled prompt's vocabulary digest, so a loader can + refuse a front-end that does not pair with its ``prompt.json``. + plan_hash: the compiled prompt's plan digest. + bundle_uuid: identity of the SID bundle the prompt was compiled against. + """ + + embed_keys: List[List[str]] + vocab_hash: str + plan_hash: str + bundle_uuid: str + + def __init__( + self, + assembler: PromptAssembler, + projections: List[nn.Module], + embed_keys: List[List[str]], + tables: Optional[List[SlotTable]] = None, + vocab_hash: str = "", + plan_hash: str = "", + bundle_uuid: str = "", + ) -> None: + super().__init__() + self.assembler = assembler + self.projections = nn.ModuleList(projections) + self.embed_keys = embed_keys + self.has_tables = tables is not None + self.tables = nn.ModuleList(tables if tables is not None else []) + self.vocab_hash = vocab_hash + self.plan_hash = plan_hash + self.bundle_uuid = bundle_uuid + + def forward( + self, + data: Dict[str, torch.Tensor], + device: Optional[torch.device] = None, + ) -> Dict[str, torch.Tensor]: + """Assemble one batch and project its holes. + + The second argument is not decoration: a C++ host that runs this as its + JIT stage calls ``forward(data, device)``, the same pair tzrec's own + ScriptWrapper takes. A Python caller may omit it, in which case the + batch stays where it is. + + Args: + data: the parsed feature dict, plus looked-up rows under a host + lookup stage. + device: where to run; the batch is moved there first. + + Returns: + The assembler's outputs plus ``slot_embeds``. + """ + target = _pick_device(data) if device is None else device + batch: Dict[str, torch.Tensor] = {} + for key, value in data.items(): + batch[key] = value.to(target) + out = self.assembler(batch) + + # gathered first, into a plain list: TorchScript indexes a ModuleList + # only with a literal, so the two lists cannot be walked together + features: List[torch.Tensor] = [] + if self.has_tables: + for table in self.tables: + features.append(table(batch)) + else: + for keys in self.embed_keys: + members: List[torch.Tensor] = [] + for key in keys: + members.append(batch[key]) + features.append(torch.cat(members, dim=1)) + + parts: List[torch.Tensor] = [] + index = 0 + for projection in self.projections: + parts.append(projection(features[index])) + index += 1 + + if len(parts) > 0: + out["slot_embeds"] = torch.cat(parts, dim=0) + else: + out["slot_embeds"] = torch.zeros( + 0, 0, dtype=torch.float32, device=out["input_ids"].device + ) + return out diff --git a/tzrec/prompt/frontend_test.py b/tzrec/prompt/frontend_test.py new file mode 100644 index 00000000..f771cc42 --- /dev/null +++ b/tzrec/prompt/frontend_test.py @@ -0,0 +1,420 @@ +# 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 unittest + +import numpy as np +import torch + +from tzrec.prompt import frontend +from tzrec.prompt.assembler import PromptAssembler as ReferenceAssembler +from tzrec.prompt.frontend import PromptAssembler, PromptFrontEnd, SlotTable, mix64 +from tzrec.prompt.types import ( + FillMode, + PromptPlan, + ResolvedSidSpace, + SlotSeg, + Static, + Width, + WidthKind, +) +from tzrec.protos.model_pb2 import FeatureGroupType +from tzrec.utils.test_util import make_test_dir + +_BASE_VOCAB = 1000 +_CODEBOOK = (4, 4, 4) +_LEVEL_OFFSETS = (0, 4, 8) + + +def _sid_space() -> ResolvedSidSpace: + return ResolvedSidSpace( + codebook=_CODEBOOK, + num_levels=3, + base_vocab_size=_BASE_VOCAB, + level_offsets=_LEVEL_OFFSETS, + band_lo=tuple(_BASE_VOCAB + o for o in _LEVEL_OFFSETS), + band_hi=tuple( + _BASE_VOCAB + o + c - 1 for o, c in zip(_LEVEL_OFFSETS, _CODEBOOK) + ), + target_vocab_size=_BASE_VOCAB + 13, + sentinel_token_id=_BASE_VOCAB + 12, + eos_token_id=2, + pad_token_id=3, + bundle_uuid="test-bundle", + ) + + +def _slot(slot_id, name, fill, sequence=True, members=None, width=None): + return SlotSeg( + slot_id=slot_id, + name=name, + feature_names=tuple(members or (name,)), + group_type=( + FeatureGroupType.JAGGED_SEQUENCE if sequence else FeatureGroupType.DEEP + ), + output_key=".sequence" if sequence else "", + fill=fill, + width=width + or (Width(WidthKind.BOUNDED, 30) if sequence else Width(WidthKind.STATIC, 1)), + ) + + +def _plan(segments, response_segments=(), projected=()): + return PromptPlan( + segments=segments, + response_segments=response_segments, + max_length=256, + max_total_length=None, + max_holes=30, + logits_suffix_len=4, + static_prefix_len=2, + projected_slots=projected, + slot_index={seg.name: i for i, seg in enumerate(projected)}, + ) + + +def _tensors(raw): + return {key: torch.from_numpy(np.asarray(value)) for key, value in raw.items()} + + +class MixTest(unittest.TestCase): + def test_mix64_matches_a_host_reference(self): + """The mixer's masked shifts must reproduce SplitMix64 on negatives.""" + + def reference(value: int) -> int: + mask = (1 << 64) - 1 + z = (value * 0x9E3779B97F4A7C15) & mask + z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & mask + z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & mask + z = z ^ (z >> 31) + return z - (1 << 64) if z >= 1 << 63 else z + + values = [0, 1, -1, 2**31, -(2**31), 123456789, -987654321] + got = mix64(torch.tensor(values, dtype=torch.int64)).tolist() + self.assertEqual(got, [reference(v) for v in values]) + + def test_output_keys_match_the_module_constants(self): + """``forward`` writes literals; they must equal the exported names.""" + module = PromptAssembler(_plan((Static((7,)),))) + keys = set(module({"batch_size": torch.tensor(2)}).keys()) + self.assertEqual( + keys, + { + frontend.OUT_INPUT_IDS, + frontend.OUT_CU_SEQLENS, + frontend.OUT_HOLE_POSITIONS, + frontend.OUT_HOLE_KEYS, + frontend.OUT_HOLE_SLOT_COUNTS, + frontend.OUT_RESPONSE_LENGTHS, + }, + ) + + +class WalkTest(unittest.TestCase): + def setUp(self): + self.hist = _slot(0, "hist", FillMode.INLINE) + self.beh = _slot(1, "beh", FillMode.PROJECTED) + self.answer = _slot( + 2, "answer", FillMode.INLINE, width=Width(WidthKind.STATIC, 3) + ) + self.plan = _plan( + segments=(Static((10, 11)), self.hist, Static((12,)), self.beh), + response_segments=(self.answer,), + projected=(self.beh,), + ) + # a SID history as a multi-value sequence: items per row, codes per item + self.raw = { + "hist.values": np.array([0, 4, 8, 1, 5, 9, 2, 6, 10], dtype=np.int64), + "hist.lengths": np.array([2, 1], dtype=np.int64), + "hist.key_lengths": np.array([3, 3, 3], dtype=np.int64), + "beh.values": np.array([7, 8, 9, 21, 22], dtype=np.int64), + "beh.lengths": np.array([3, 2], dtype=np.int64), + "answer.values": np.array([3, 7, 11, 0, 4, 8], dtype=np.int64), + "answer.lengths": np.array([1, 1], dtype=np.int64), + "answer.key_lengths": np.array([3, 3], dtype=np.int64), + } + self.module = PromptAssembler( + self.plan, + _sid_space(), + features_are_multi_valued={"hist": True, "answer": True}, + plan_hash="a1b2c3d4e5f60718", + ) + + def _assert_matches_reference(self, module, raw): + out = module(_tensors(raw)) + reference = ReferenceAssembler(self.plan, _sid_space()).forward(raw) + for key, reference_key in ( + ("input_ids", "prompt_input_ids"), + ("cu_seqlens", "prompt_cu_seqlens"), + ("hole_positions", "prompt_hole_positions"), + ("response_lengths", "prompt_response_lengths"), + ): + self.assertEqual(out[key].tolist(), reference[reference_key].tolist(), key) + return out + + def test_walk_matches_the_reference_implementation(self): + """The tensor walk and the collator's walk are one specification.""" + self._assert_matches_reference(self.module, self.raw) + + def test_walk_matches_the_reference_on_the_flat_layout(self): + """A history stored one code per position walks the same way.""" + raw = { + "hist.values": self.raw["hist.values"], + "hist.lengths": np.array([6, 3], dtype=np.int64), + "beh.values": self.raw["beh.values"], + "beh.lengths": self.raw["beh.lengths"], + "answer.values": self.raw["answer.values"], + "answer.lengths": np.array([3, 3], dtype=np.int64), + } + module = PromptAssembler(self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") + self._assert_matches_reference(module, raw) + + def test_holes_land_on_sentinels(self): + """Every recorded hole is a sentinel and every sentinel is recorded.""" + out = self.module(_tensors(self.raw)) + sentinel = _sid_space().sentinel_token_id + holes = out["hole_positions"] + self.assertTrue(bool(torch.all(out["input_ids"][holes] == sentinel))) + self.assertEqual( + int(torch.sum(out["input_ids"] == sentinel)), int(holes.numel()) + ) + self.assertEqual(out["hole_slot_counts"].tolist(), [5]) + + def test_scripting_preserves_every_output(self): + """The artifact and the eager module are the same function.""" + batch = _tensors(self.raw) + eager = self.module(batch) + scripted = torch.jit.script(self.module)(batch) + for key, value in eager.items(): + self.assertTrue(torch.equal(scripted[key], value), key) + + def test_inline_without_a_sid_space_is_rejected(self): + with self.assertRaisesRegex(ValueError, "no sid_space"): + PromptAssembler(self.plan) + + +class FoldTest(unittest.TestCase): + def setUp(self): + self.beh = _slot(0, "beh", FillMode.PROJECTED) + self.plan = _plan(segments=(self.beh,), projected=(self.beh,)) + self.module = PromptAssembler( + self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718" + ) + + def _keys(self, values, lengths): + return self.module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array(lengths, dtype=np.int64), + } + ) + )["hole_keys"] + + def test_equal_content_folds_equal(self): + """A key is a function of the hole's inputs and nothing else.""" + self.assertTrue( + torch.equal(self._keys([5, 6, 7], [3]), self._keys([5, 6, 7], [3])) + ) + + def test_different_content_folds_apart(self): + """The whole point: a changed input must not reuse the cached KV.""" + first = self._keys([5, 6, 7], [3]) + second = self._keys([5, 6, 8], [3]) + self.assertEqual(first[:2].tolist(), second[:2].tolist()) + self.assertNotEqual(int(first[2]), int(second[2])) + + def test_the_plan_hash_salts_the_keys(self): + """Two plans must not cross-match in a shared prefix cache.""" + other = PromptAssembler(self.plan, _sid_space(), plan_hash="ffffffffffffffff") + batch = _tensors( + { + "beh.values": np.array([5, 6, 7], dtype=np.int64), + "beh.lengths": np.array([3], dtype=np.int64), + } + ) + self.assertFalse( + torch.equal(self.module(batch)["hole_keys"], other(batch)["hole_keys"]) + ) + + def test_a_permuted_multi_value_item_does_not_collide(self): + """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" + module = PromptAssembler( + self.plan, + _sid_space(), + features_are_multi_valued={"beh": True}, + plan_hash="a1b2c3d4e5f60718", + ) + + def keys(values): + return module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array([1], dtype=np.int64), + "beh.key_lengths": np.array([3], dtype=np.int64), + } + ) + )["hole_keys"] + + self.assertNotEqual(keys([1, 5, 9]).tolist(), keys([9, 5, 1]).tolist()) + + def test_two_members_exchanging_values_do_not_collide(self): + """Without the member index a two-member slot is order-blind.""" + slot = _slot(0, "pair", FillMode.PROJECTED, members=("a", "b")) + plan = _plan(segments=(slot,), projected=(slot,)) + module = PromptAssembler(plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") + + def keys(first, second): + return module( + _tensors( + { + "a.values": np.array(first, dtype=np.int64), + "a.lengths": np.array([1], dtype=np.int64), + "b.values": np.array(second, dtype=np.int64), + "b.lengths": np.array([1], dtype=np.int64), + } + ) + )["hole_keys"] + + self.assertNotEqual(keys([3], [9]).tolist(), keys([9], [3]).tolist()) + + def test_a_dense_member_folds_its_bit_pattern(self): + """A float member contributes the parsed input verbatim, so it is stable.""" + slot = _slot(0, "vec", FillMode.PROJECTED, sequence=False) + plan = _plan(segments=(slot,), projected=(slot,)) + module = PromptAssembler( + plan, _sid_space(), features_are_dense={"vec": True}, plan_hash="a1b2" + ) + + def keys(rows): + return module( + { + "vec.values": torch.tensor(rows, dtype=torch.float32), + "vec.lengths": torch.ones(len(rows), dtype=torch.int64), + } + )["hole_keys"] + + self.assertEqual(keys([[0.5, 1.0]]).tolist(), keys([[0.5, 1.0]]).tolist()) + self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) + + @unittest.skipIf(not torch.cuda.is_available(), "no GPU") + def test_the_fold_is_bit_identical_across_devices(self): + """Integer addition cannot depend on the order a device reduces in.""" + batch = _tensors( + { + "beh.values": np.arange(64, dtype=np.int64), + "beh.lengths": np.array([32, 32], dtype=np.int64), + } + ) + on_cpu = self.module(batch)["hole_keys"] + on_gpu = self.module({k: v.cuda() for k, v in batch.items()})["hole_keys"] + self.assertTrue(torch.equal(on_cpu, on_gpu.cpu())) + + +class FrontEndTest(unittest.TestCase): + def setUp(self): + self.slot = _slot(0, "beh", FillMode.PROJECTED) + self.plan = _plan(segments=(Static((10,)), self.slot), projected=(self.slot,)) + self.batch = _tensors( + { + "beh.values": np.array([1, 2, 3], dtype=np.int64), + "beh.lengths": np.array([3], dtype=np.int64), + } + ) + + def _table(self): + torch.manual_seed(0) + return SlotTable(["beh.values"], ["beh.lengths"], [""], [16], [8], True) + + def _front_end(self, tables, embed_keys=(("beh",),)): + torch.manual_seed(1) + assembler = PromptAssembler(self.plan, _sid_space(), plan_hash="a1b2") + projection = torch.nn.Linear(8, 16) + return PromptFrontEnd( + assembler, + [projection], + [list(keys) for keys in embed_keys], + tables=tables, + vocab_hash="vocab", + plan_hash="plan", + bundle_uuid="test-bundle", + ) + + def test_in_module_tables_produce_one_embedding_per_hole(self): + """Without a host lookup stage, the artifact carries the tables.""" + out = self._front_end([self._table()])(self.batch) + self.assertEqual(tuple(out["slot_embeds"].shape), (3, 16)) + self.assertEqual(int(out["hole_positions"].numel()), 3) + + def test_both_lookup_shapes_agree_given_the_same_rows(self): + """Where the lookup happens must not change what the model sees.""" + table = self._table() + with_tables = self._front_end([table]) + host = self._front_end(None) + # a host stage hands over rows and item counts, not ids and .lengths + batch = { + "beh.values": self.batch["beh.values"], + "beh": table(self.batch), + frontend.host_lengths_key("beh"): self.batch["beh.lengths"], + } + self.assertTrue( + torch.allclose( + with_tables(self.batch)["slot_embeds"], host(batch)["slot_embeds"] + ) + ) + self.assertTrue( + torch.equal(with_tables(self.batch)["hole_keys"], host(batch)["hole_keys"]) + ) + + def test_the_device_argument_is_optional(self): + """The processor passes a device like ScriptWrapper; sglang does not.""" + module = torch.jit.script(self._front_end([self._table()]).eval()) + implicit = module(self.batch) + explicit = module(self.batch, torch.device("cpu")) + for key, value in implicit.items(): + self.assertTrue(torch.equal(explicit[key], value), key) + + def test_the_artifact_reloads_with_its_identity(self): + """A loader pairs the front-end with prompt.json by these attributes.""" + path = os.path.join(make_test_dir(), "scripted_model.pt") + torch.jit.script(self._front_end([self._table()]).eval()).save(path) + loaded = torch.jit.load(path) + self.assertEqual( + (loaded.vocab_hash, loaded.plan_hash, loaded.bundle_uuid), + ("vocab", "plan", "test-bundle"), + ) + out = loaded(self.batch) + self.assertEqual(tuple(out["slot_embeds"].shape), (3, 16)) + + def test_pattern_i_exports_empty_hole_streams(self): + """With no projected slot the artifact degenerates, and still exists.""" + hist = _slot(0, "hist", FillMode.INLINE) + plan = _plan(segments=(Static((10,)), hist)) + module = torch.jit.script( + PromptFrontEnd(PromptAssembler(plan, _sid_space()), [], []).eval() + ) + out = module( + _tensors( + { + "hist.values": np.array([0, 5, 10], dtype=np.int64), + "hist.lengths": np.array([3], dtype=np.int64), + } + ) + ) + self.assertEqual(out["input_ids"].tolist(), [10, 1000, 1005, 1010]) + self.assertEqual(int(out["hole_positions"].numel()), 0) + self.assertEqual(tuple(out["slot_embeds"].shape), (0, 0)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py index a152c8e8..81b003af 100644 --- a/tzrec/prompt/persist.py +++ b/tzrec/prompt/persist.py @@ -9,20 +9,92 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Checks a checkpoint's prompt contract on restore. +"""Checks a checkpoint's prompt contract on restore, and publishes it on export. The digests ride in the HF export metadata the checkpoint already carries, so -there is no second file that can drift from the weights beside it. +there is no second file that can drift from the weights beside it. Export +additionally writes ``prompt/prompt.json``: what a serving runtime reads, and +only that. The plan itself is deliberately not published there -- it reaches +serving compiled into the front-end artifact, and a copy a runtime could +interpret would invite the second assembler this design exists to prevent. """ +import dataclasses import json import os -from typing import Dict, Optional +from typing import Any, Dict, List, Optional from tzrec.constant import HF_EXPORT_META_FILENAME -from tzrec.prompt.types import CompiledPrompt +from tzrec.prompt.types import CompiledPrompt, SlotSeg, Static from tzrec.utils.logging_util import logger +PROMPT_DIR = "prompt" +PROMPT_CONTRACT_FILENAME = "prompt.json" +TOKENIZER_DIR = "tokenizer" + + +def _decode_schedule(compiled_prompt: CompiledPrompt) -> List[Dict[str, int]]: + """One entry per response position: a forced token or the band it is masked to. + + Derived from the response segments, so a runtime gets the schedule without + gaining the ability to interpret a plan. + """ + sid_space = compiled_prompt.sid_space + schedule: List[Dict[str, int]] = [] + for seg in compiled_prompt.prompt_plan.response_segments: + if isinstance(seg, Static): + schedule.extend({"token_id": int(t)} for t in seg.token_ids) + continue + assert isinstance(seg, SlotSeg) and sid_space is not None + width = seg.width.num_positions or 0 + for index in range(width): + level = index % sid_space.num_levels + schedule.append( + { + "level": level, + "band_lo": sid_space.band_lo[level], + "band_hi": sid_space.band_hi[level], + } + ) + return schedule + + +def write_serving_contract( + compiled_prompt: CompiledPrompt, export_dir: str, frontend: Dict[str, Any] +) -> str: + """Write ``prompt/prompt.json``, everything a serving runtime reads. + + Args: + compiled_prompt: the compiled prompt. + export_dir: the HuggingFace export directory. + frontend: how the exported front-end is laid out and what it reads. + + Returns: + The path written. + """ + plan = compiled_prompt.prompt_plan + out = os.path.join(export_dir, PROMPT_DIR) + os.makedirs(out, exist_ok=True) + path = os.path.join(out, PROMPT_CONTRACT_FILENAME) + with open(path, "w") as f: + json.dump( + { + "sid_space": dataclasses.asdict(compiled_prompt.sid_space) + if compiled_prompt.sid_space is not None + else None, + "decode": _decode_schedule(compiled_prompt), + "static_prefix_len": plan.static_prefix_len, + "max_length": plan.max_length, + "max_total_length": plan.max_total_length, + "vocab_hash": compiled_prompt.vocab_hash, + "plan_hash": compiled_prompt.plan_hash, + "frontend": frontend, + }, + f, + indent=2, + ) + return path + def read_prompt_digests(source_dir: str) -> Optional[Dict[str, str]]: """Read the digests a checkpoint recorded, or None when it has none.""" diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index 7cb65226..efdd7256 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -16,7 +16,7 @@ physical dimension: the model resolves those at ``__init__``. """ -from dataclasses import dataclass +from dataclasses import dataclass, field from enum import Enum from typing import Mapping, Optional, Tuple, Union @@ -84,6 +84,10 @@ class ResolvedSidSpace: slot is projected. eos_token_id: end-of-sequence id of the extended tokenizer. pad_token_id: padding id of the extended tokenizer. + bundle_uuid: identity of the SID bundle this space was compiled + against, empty when no manifest was read. Serving refuses a + catalog whose bundle differs: a copied artifact's path proves + nothing. """ codebook: Tuple[int, ...] @@ -96,6 +100,7 @@ class ResolvedSidSpace: sentinel_token_id: Optional[int] eos_token_id: int pad_token_id: int + bundle_uuid: str = "" @dataclass(frozen=True) @@ -135,6 +140,36 @@ class SlotSeg: Segment = Union[Static, SlotSeg] +@dataclass(frozen=True) +class FoldConstants: + """Odd multipliers mixed into ``hole_keys``. + + Written as signed int64 so torch takes them verbatim: the fold wraps, and a + host that had to convert them would be a second place to get it wrong. + + They live in the plan rather than in the host so the artifact, not the + machine that runs it, decides the keys. Each closes a collision that would + otherwise produce a correct-looking prefix-cache hit: two slots holding the + same id, two members of one slot exchanging values, or a multi-value item + permuted. + + Args: + slot: multiplies ``slot_id`` into the per-hole salt. + plan: multiplies the low 64 bits of ``plan_hash``, so keys are + artifact-specific and a rolling upgrade cannot cross-match. + value: multiplies each contributing value. + index: multiplies the member-and-position index within a hole. + position: multiplies the hole index in the per-item outer fold, without + which a permuted history collides. + """ + + slot: int = -7046029254386353131 + plan: int = -4417276706812531889 + value: int = -49064778989728563 + index: int = -2960836687051489901 + position: int = -6752110988234923001 + + @dataclass(frozen=True) class PromptPlan: """The walk order the assembler follows, plus the ceilings derived from it. @@ -147,7 +182,11 @@ class PromptPlan: max_holes: per-row projected-position ceiling, not a runtime shape. logits_suffix_len: upper bound on the supervised logits window. static_prefix_len: leading positions that are request-invariant. - projected_slots: PROJECTED occurrences in emission order. + projected_slots: PROJECTED occurrences in emission order, which is also + ascending hole position; nothing may reorder them by slot id or by + shared module, because the serving scatter is positional. + slot_index: slot name to its index in ``projected_slots``. + fold: the constants ``hole_keys`` mixes in. """ segments: Tuple[Segment, ...] @@ -158,6 +197,8 @@ class PromptPlan: logits_suffix_len: Optional[int] static_prefix_len: int projected_slots: Tuple[SlotSeg, ...] + slot_index: Mapping[str, int] = field(default_factory=dict) + fold: FoldConstants = field(default_factory=FoldConstants) @dataclass(frozen=True) diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py new file mode 100644 index 00000000..ff97211d --- /dev/null +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -0,0 +1,178 @@ +# 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. + +"""Reads a genrec export the way the SGLang genrec stack does. + +No sglang import: these tests pin the contract from the consumer's side -- +the composite config its model wrapper resolves, the front-end outputs its +multimodal processor cuts into items, the host-side item hash it re-folds +from ``hole_keys``, and the ``prompt.json`` fields its constraint-index +builder reads -- so a change on this side that would break serving fails here. +""" + +import io +import json +import os +import unittest + +import numpy as np +import torch +from safetensors.torch import load_file + +from tzrec.prompt.frontend import mix64 +from tzrec.prompt.types import FoldConstants +from tzrec.tests.prompt_test_util import export_tiny_genrec, offset_sid_codes +from tzrec.utils.test_util import make_test_dir + +# sglang's multimodal/processors/prompt_genrec.py folds a slot's per-hole keys +# into one item hash with this mixer and this position constant +_C_POSITION = -6752110988234923001 +_MASK64 = (1 << 64) - 1 + + +def _mix64_host(value: int) -> int: + z = (value * 0x9E3779B97F4A7C15) & _MASK64 + z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & _MASK64 + z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & _MASK64 + return z ^ (z >> 31) + + +def _item_hash(keys) -> int: + total = 0 + for index, key in enumerate(keys.tolist()): + total = (total + _mix64_host((key + _C_POSITION * index) & _MASK64)) & _MASK64 + return total + + +class GenrecServingContractTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + cls.exported = export_tiny_genrec(make_test_dir(), bundle_uuid="bundle-test") + with open(os.path.join(cls.exported.export_dir, "config.json"), "r") as f: + cls.config = json.load(f) + with open( + os.path.join(cls.exported.export_dir, "prompt", "prompt.json"), "r" + ) as f: + cls.contract = json.load(f) + cls.front_end = torch.jit.load( + os.path.join(cls.exported.export_dir, "frontend", "scripted_model.pt") + ) + + def _payload(self): + """A request as the in-process processor receives it: an npz blob.""" + buffer = io.BytesIO() + np.savez_compressed( + buffer, + **{ + "hist.values": offset_sid_codes([0, 1, 2, 3, 0, 1], [4, 4, 4]), + "hist.lengths": np.array([6], dtype=np.int64), + "beh.values": np.array([3, 9], dtype=np.int64), + "beh.lengths": np.array([2], dtype=np.int64), + }, + ) + buffer.seek(0) + with np.load(buffer, allow_pickle=False) as data: + return {key: torch.from_numpy(np.asarray(data[key])) for key in data.files} + + def test_config_is_composite_and_names_the_backbone(self) -> None: + self.assertEqual(self.config["architectures"], ["PromptGenRecForCausalLM"]) + self.assertEqual(self.config["model_type"], "prompt_genrec") + text_config = self.config["text_config"] + self.assertEqual(text_config["model_type"], "qwen2") + self.assertEqual( + text_config["vocab_size"], + self.exported.compiled_prompt.sid_space.target_vocab_size, + ) + + def test_weights_keep_the_backbones_own_names(self) -> None: + """The wrapper delegates load_weights wholesale, so no remapping exists.""" + from transformers import AutoConfig, AutoModelForCausalLM + + config = AutoConfig.for_model(**self.config["text_config"]) + with torch.device("meta"): + backbone = AutoModelForCausalLM.from_config(config) + expected = set(backbone.state_dict().keys()) + exported = set( + load_file(os.path.join(self.exported.export_dir, "model.safetensors")) + ) + self.assertTrue(exported <= expected, exported - expected) + self.assertIn("model.embed_tokens.weight", exported) + + def test_front_end_output_feeds_the_processor(self) -> None: + out = self.front_end(self._payload()) + for key in ( + "input_ids", + "hole_positions", + "slot_embeds", + "hole_keys", + "hole_slot_counts", + ): + self.assertIn(key, out) + positions = out["hole_positions"] + self.assertEqual(positions.dtype, torch.int64) + self.assertTrue(bool(torch.all(positions[1:] > positions[:-1]))) + self.assertEqual(int(out["hole_slot_counts"].sum()), int(positions.numel())) + self.assertEqual( + tuple(out["slot_embeds"].shape), + (int(positions.numel()), self.config["text_config"]["hidden_size"]), + ) + self.assertEqual(out["slot_embeds"].dtype, torch.float32) + sentinel = self.contract["sid_space"]["sentinel_token_id"] + self.assertTrue(bool(torch.all(out["input_ids"][positions] == sentinel))) + self.assertEqual(out["hole_keys"].dtype, torch.int64) + + def test_host_item_hash_refolds_hole_keys(self) -> None: + """The per-item outer fold on the host uses the plan's position constant.""" + self.assertEqual(FoldConstants().position, _C_POSITION) + for value in (0, 1, -1, 123456789, -(2**40)): + signed = _mix64_host(value & _MASK64) + signed = signed - (1 << 64) if signed >= 1 << 63 else signed + self.assertEqual(int(mix64(torch.tensor([value]))[0]), signed) + keys = self.front_end(self._payload())["hole_keys"] + self.assertEqual(_item_hash(keys), _item_hash(keys)) + self.assertNotEqual(_item_hash(keys), _item_hash(keys.flip(0))) + + def test_prompt_json_carries_the_index_builder_inputs(self) -> None: + space = self.contract["sid_space"] + compiled = self.exported.compiled_prompt.sid_space + self.assertEqual(space["base_vocab_size"], compiled.base_vocab_size) + self.assertEqual(space["num_levels"], 3) + self.assertEqual(space["band_lo"], list(compiled.band_lo)) + self.assertEqual(space["band_hi"], list(compiled.band_hi)) + self.assertEqual(space["bundle_uuid"], "bundle-test") + self.assertEqual(len(self.contract["decode"]), 3) + self.assertEqual( + self.contract["vocab_hash"], self.exported.compiled_prompt.vocab_hash + ) + self.assertEqual(self.contract["frontend"]["dir"], "frontend") + self.assertEqual(self.contract["frontend"]["model"], "scripted_model.pt") + # the artifact pairs with the contract by identity, not by path + self.assertEqual(self.front_end.vocab_hash, self.contract["vocab_hash"]) + self.assertEqual(self.front_end.plan_hash, self.contract["plan_hash"]) + self.assertEqual(self.front_end.bundle_uuid, "bundle-test") + + def test_tokenizer_dir_decodes_a_sid_atom(self) -> None: + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained( + os.path.join(self.exported.export_dir, "prompt", "tokenizer") + ) + space = self.contract["sid_space"] + self.assertEqual(tokenizer.decode([space["band_lo"][0]]), "<|sid_0|>") + self.assertEqual( + tokenizer.convert_ids_to_tokens(space["sentinel_token_id"]), "<|pg_hole|>" + ) + self.assertEqual(tokenizer.eos_token_id, space["eos_token_id"]) + self.assertEqual(tokenizer.pad_token_id, space["pad_token_id"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index ed86e8a0..42278a42 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -9,20 +9,26 @@ # See the License for the specific language governing permissions and # limitations under the License. +import json import os import unittest +from unittest import mock import numpy as np import torch import torch.fx from tzrec.datasets.utils import Batch +from tzrec.main import export from tzrec.models.model import TrainWrapper +from tzrec.prompt.assembler import PromptAssembler from tzrec.tests.prompt_test_util import ( GenrecModelTestBase, assemble_into, + export_tiny_genrec, offset_sid_codes, ) +from tzrec.utils.test_util import make_test_dir _CODEBOOK = [4, 4, 4] _WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] @@ -80,5 +86,111 @@ def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(model)) +def _serving_batch(): + """One request as the front-end reads it: an INLINE history, one behaviour.""" + return { + "hist.values": torch.tensor(offset_sid_codes([0, 1, 2, 3, 0, 1], _CODEBOOK)), + "hist.lengths": torch.tensor([6]), + "beh.values": torch.tensor([3, 9]), + "beh.lengths": torch.tensor([2]), + } + + +class GenrecExportIntegrationTest(unittest.TestCase): + """checkpoint -> export -> the artifacts a serving runtime loads.""" + + def setUp(self) -> None: + self.test_dir = make_test_dir() + + def test_export_writes_a_loadable_serving_directory(self) -> None: + exported = export_tiny_genrec(self.test_dir) + for name in ( + "config.json", + "model.safetensors", + "prompt/prompt.json", + "prompt/tokenizer/tokenizer.json", + "prompt/tokenizer/tokenizer_config.json", + "frontend/scripted_model.pt", + "frontend/fg.json", + "frontend/pipeline.config", + "frontend/model_acc.json", + ): + self.assertTrue( + os.path.exists(os.path.join(exported.export_dir, name)), name + ) + self.assertFalse( + os.path.exists(os.path.join(exported.export_dir, "frontend/sparse")) + ) + + # the collator and the scripted front-end are two call sites of one walk + batch = _serving_batch() + compiled = exported.compiled_prompt + collator = PromptAssembler( + compiled.prompt_plan, compiled.sid_space, include_response=False + ).forward({k: v.numpy() for k, v in batch.items()}) + front_end = torch.jit.load( + os.path.join(exported.export_dir, "frontend", "scripted_model.pt") + ) + out = front_end(batch, torch.device("cpu")) + self.assertEqual( + out["input_ids"].tolist(), collator["prompt_input_ids"].tolist() + ) + self.assertEqual( + out["cu_seqlens"].tolist(), collator["prompt_cu_seqlens"].tolist() + ) + self.assertEqual( + out["hole_positions"].tolist(), collator["prompt_hole_positions"].tolist() + ) + self.assertEqual(tuple(out["slot_embeds"].shape), (2, 32)) + + def test_distributed_embedding_export_writes_the_processor_shape(self) -> None: + exported = export_tiny_genrec(self.test_dir) + dist_dir = os.path.join(self.test_dir, "export_dist") + with mock.patch.dict(os.environ, {"USE_DISTRIBUTED_EMBEDDING": "1"}): + export( + exported.config_path, dist_dir, checkpoint_path=exported.checkpoint_dir + ) + frontend_dir = os.path.join(dist_dir, "frontend") + sparse_dir = os.path.join(frontend_dir, "sparse") + for name in ( + "sparse_embeddings-00-of-01.npz", + "sparse_embeddings-00-of-01.json", + "sparse_embedding.json", + "sparse_features.json", + ): + self.assertTrue(os.path.exists(os.path.join(sparse_dir, name)), name) + with open(os.path.join(frontend_dir, "model_acc.json"), "r") as f: + acc = json.load(f) + self.assertEqual(acc["DISTRIBUTED_EMBEDDING"], "1") + self.assertEqual(acc["INPUT_TILE"], "3") + with open(os.path.join(frontend_dir, "dense_meta.json"), "r") as f: + dense_meta = json.load(f) + self.assertEqual(dense_meta, {"sequence__ec": ["beh__ec", "beh__lengths"]}) + with open(os.path.join(dist_dir, "prompt", "prompt.json"), "r") as f: + self.assertEqual(json.load(f)["frontend"]["lookup"], "host") + + # a processor simulator: look the ids up in the exported tables the way + # the distributed-embedding stage does, then feed the dense stage + batch = _serving_batch() + with open(os.path.join(sparse_dir, "sparse_features.json"), "r") as f: + table_name = json.load(f)["beh__ec"]["embedding_name"] + with np.load(os.path.join(sparse_dir, "sparse_embeddings-00-of-01.npz")) as npz: + table = npz[table_name] + host = dict(batch) + host["beh"] = torch.from_numpy( + table[batch["beh.values"].numpy()].astype(np.float32) + ) + host["beh__lengths"] = batch["beh.lengths"] + got = torch.jit.load(os.path.join(frontend_dir, "scripted_model.pt"))(host) + expected = torch.jit.load( + os.path.join(exported.export_dir, "frontend", "scripted_model.pt") + )(batch) + self.assertTrue( + torch.allclose(got["slot_embeds"], expected["slot_embeds"], atol=1e-6) + ) + self.assertTrue(torch.equal(got["hole_keys"], expected["hole_keys"])) + self.assertEqual(got["input_ids"].tolist(), expected["input_ids"].tolist()) + + if __name__ == "__main__": unittest.main() diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index 62da6a07..396aca96 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -9,9 +9,11 @@ # See the License for the specific language governing permissions and # limitations under the License. +import dataclasses +import json import os import unittest -from typing import Any, Dict, Sequence +from typing import Any, Dict, List, Optional, Sequence import numpy as np import torch @@ -20,13 +22,17 @@ from tzrec.datasets.utils import BASE_DATA_GROUP, Batch from tzrec.features.feature import BaseFeature, FgMode, create_features -from tzrec.main import _create_model +from tzrec.main import _create_features, _create_model, export +from tzrec.models.model import TrainWrapper from tzrec.prompt.assembler import PromptAssembler from tzrec.prompt.compile import compile_prompt from tzrec.prompt.types import CompiledPrompt from tzrec.protos import feature_pb2 from tzrec.protos.model_pb2 import ModelConfig +from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig +from tzrec.utils.hf_export_util import write_hf_assets +from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import create_tiny_causal_lm, make_test_dir @@ -163,3 +169,133 @@ def _batch(self, parsed, compiled_prompt=None, sparse=None): {k: torch.from_numpy(np.asarray(v)) for k, v in streams.items()} ) return batch + + +def write_genrec_checkpoint(model: torch.nn.Module, ckpt_dir: str) -> str: + """Save a model the way a training run does, minus the dynamic-table dump. + + Args: + model: the bare genrec model; wrapped and initialized here. + ckpt_dir: the ``model.ckpt-N`` directory to write. + + Returns: + ``ckpt_dir``. + """ + from torch.distributed.checkpoint import save + + wrapped = TrainWrapper(model) + init_parameters(wrapped, torch.device("cpu")) + save(wrapped.state_dict(), checkpoint_id=os.path.join(ckpt_dir, "model")) + write_hf_assets(wrapped, ckpt_dir) + return ckpt_dir + + +@dataclasses.dataclass +class ExportedGenrec: + """A tiny genrec model, its checkpoint and its HF export. + + Args: + config: the pipeline config the export ran on. + config_path: where it was written. + features: the created features. + compiled_prompt: the prompt as the export compiled it. + checkpoint_dir: the checkpoint the export converted. + export_dir: the export directory. + """ + + config: EasyRecConfig + config_path: str + features: List[BaseFeature] + compiled_prompt: CompiledPrompt + checkpoint_dir: str + export_dir: str + + +def export_tiny_genrec( + test_dir: str, + projected: bool = True, + bundle_uuid: Optional[str] = None, + hidden_size: int = 32, +) -> ExportedGenrec: + """Train nothing, save a checkpoint of a tiny genrec model, and export it. + + The prompt carries an INLINE SID history and, when ``projected``, one + PROJECTED behaviour slot, which is what makes the export write a front-end + with tables and a projection. + + Args: + test_dir: scratch directory. + projected: whether to add the projected slot. + bundle_uuid: when given, a SID manifest carrying it is written and + referenced, so the export records a bundle identity. + hidden_size: the tiny backbone's hidden size. + + Returns: + Everything a test needs to read the export back. + """ + backbone = os.path.join(test_dir, "backbone") + create_tiny_causal_lm(64).save_pretrained(backbone) + tok = create_prompt_tokenizer(os.path.join(test_dir, "tok.json"), _WORDS) + manifest = "" + if bundle_uuid is not None: + manifest = os.path.join(test_dir, "manifest.json") + with open(manifest, "w") as f: + json.dump({"codebook": _CODEBOOK, "bundle_uuid": bundle_uuid}, f) + + config = EasyRecConfig() + text_format.Merge( + f''' +train_input_path: "" eval_input_path: "" model_dir: "{test_dir}/train" +train_config {{ + sparse_optimizer {{ adagrad_optimizer {{ lr: 0.0 }} constant_learning_rate {{}} }} + dense_optimizer {{ adam_optimizer {{ lr: 0.0001 }} constant_learning_rate {{}} }} + num_epochs: 1 +}} +data_config {{ + batch_size: 4 dataset_type: ParquetDataset fg_mode: FG_NONE + label_fields: "answer" num_workers: 1 +}} +feature_configs {{ {_HIST} }} +{"feature_configs { " + projected_feature("beh", 8) + " }" if projected else ""} +prompt_config {{ + tokenizer_path: "{tok}" + prompt: "History : {{{{hist}}}} .{" {{beh}}" if projected else ""} Predict :" + response: "{{{{answer}}}}" + sid_space {{ + codebook: 4 codebook: 4 codebook: 4 + {f'manifest_path: "{manifest}"' if manifest else ""} + }} + max_length: 64 +}} +model_config {{ + genrec_causal_lm_model {{ + hf_model_name_or_path: "{backbone}" + common {{ beam_widths: 2 beam_widths: 2 beam_widths: 2 num_return_sequences: 2 }} + }} +}} +''', + config, + ) + os.makedirs(config.model_dir) + config_path = os.path.join(test_dir, "pipeline.config") + with open(config_path, "w") as f: + f.write(text_format.MessageToString(config)) + + features = _create_features(list(config.feature_configs), config.data_config) + compiled_prompt = compile_prompt(config.prompt_config, features, ["answer"]) + model = _create_model( + config.model_config, features, ["answer"], compiled_prompt=compiled_prompt + ) + checkpoint_dir = write_genrec_checkpoint( + model, os.path.join(config.model_dir, "model.ckpt-1") + ) + export_dir = os.path.join(test_dir, "export") + export(config_path, export_dir, checkpoint_path=checkpoint_dir) + return ExportedGenrec( + config=config, + config_path=config_path, + features=features, + compiled_prompt=compiled_prompt, + checkpoint_dir=checkpoint_dir, + export_dir=export_dir, + ) diff --git a/tzrec/utils/export_util.py b/tzrec/utils/export_util.py index edb095bf..7383c2ce 100644 --- a/tzrec/utils/export_util.py +++ b/tzrec/utils/export_util.py @@ -2675,6 +2675,9 @@ def _add_sparse_table( continue table_fqn = name[: -len(".weight")] export_emb_name = checkpoint_util.remap_input_tile_user_key(table_fqn) + # a dense weight beside the tables, such as a genrec backbone's + if export_emb_name not in emb_name_to_emb_dim: + continue state_values_by_emb[export_emb_name] = values for export_emb_name, values in state_values_by_emb.items(): From b3079968df39a8e524d9bbf80a33a42e282c2d33 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 14:52:00 +0800 Subject: [PATCH 02/21] [refactor] make the prompt assembler one scripted module with two call sites RFC 0001 fixes the assembler as one TorchScript module the collator calls on the host and export wraps as the serving front-end, but the serving work added a second walk beside the collator's numpy one and pinned the two with a parity test. The tensor-op walk now lives in assembler.py as the only PromptAssembler: the collator feeds it the parsed tensors directly and prefixes its streams into additional_infos, the exported front-end scripts the same module, and the numpy walk's validation moves into it so a bad row fails identically at both call sites. Holes are indexed by projected occurrence in emission order rather than by slot name, so a slot that appears twice in the template keeps both of its hole groups, and the plan drops the slot_index that collapsed them. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/datasets/dataset.py | 8 +- tzrec/prompt/assembler.py | 827 ++++++++++++++------ tzrec/prompt/assembler_test.py | 337 ++++++-- tzrec/prompt/compile.py | 1 - tzrec/prompt/export.py | 36 +- tzrec/prompt/frontend.py | 542 +------------ tzrec/prompt/frontend_test.py | 272 +------ tzrec/prompt/types.py | 2 - tzrec/tests/genrec_serving_contract_test.py | 2 +- tzrec/tests/prompt_integration_test.py | 33 +- tzrec/tests/prompt_test_util.py | 25 +- 11 files changed, 944 insertions(+), 1141 deletions(-) diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index aced406d..33b48235 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -18,7 +18,6 @@ import numpy as np import pyarrow as pa import pyarrow.compute as pc -import torch from torch import distributed as dist from torch.utils.data import DataLoader, IterableDataset, get_worker_info @@ -41,7 +40,7 @@ remove_nullable, ) from tzrec.features.feature import BaseFeature -from tzrec.prompt.assembler import PromptAssembler +from tzrec.prompt.assembler import PROMPT_INFO_PREFIX, PromptAssembler from tzrec.prompt.types import CompiledPrompt from tzrec.protos import data_pb2 from tzrec.utils import config_util @@ -118,6 +117,7 @@ def __init__( PromptAssembler( compiled_prompt.prompt_plan, compiled_prompt.sid_space, + plan_hash=compiled_prompt.plan_hash, include_response=mode != Mode.PREDICT, ) if compiled_prompt is not None @@ -403,8 +403,8 @@ def _build_batch(self, input_data: Dict[str, pa.Array]) -> Batch: if self._prompt_assembler is not None: batch.additional_infos.update( { - k: torch.from_numpy(np.asarray(v)) - for k, v in self._prompt_assembler.forward(output_data).items() + PROMPT_INFO_PREFIX + k: v + for k, v in self._prompt_assembler(output_data).items() } ) diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index b866b87e..a8c0cd05 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -9,320 +9,649 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Builds the packed token stream a compiled prompt describes. +"""The prompt assembler: one scripted walk with two call sites. -Runs in the dataloader worker, after the features are parsed and outside any -feature's ``_parse``. Pure integer arithmetic with no FG dependency, so the -same walk is portable to an online C++/Java processor. +It runs in the dataloader worker after the features are parsed, and again +inside the exported front-end at serving. ``PromptPlan`` is a compile-time +constant, so the segment loop unrolls into parallel constant lists at +construction and what remains is jagged integer arithmetic -- ``cumsum``, +``repeat_interleave``, ``index_copy_`` -- with no data-dependent control flow, +which is what lets ``torch.jit.script`` carry the same module into a runtime +that has no tzrec source. Serving therefore never reimplements this walk. + +``hole_keys`` is computed at both call sites and discarded by training. It is +integer end to end: ``int64`` addition is associative and commutative and wraps +deterministically, so the fold cannot depend on the order ``index_add_`` +happens to reduce in, on the device, or on how the batch was split. A float +accumulator would satisfy "do not fold the projected vector" in letter and +reintroduce the variance in spirit. """ -from dataclasses import dataclass -from typing import Dict, List, Optional +import hashlib +from typing import Dict, Final, List, Optional, Tuple -import numpy as np +import torch +from torch import nn from tzrec.prompt.types import ( FillMode, + FoldConstants, PromptPlan, ResolvedSidSpace, SlotSeg, Static, ) from tzrec.protos.model_pb2 import FeatureGroupType -from tzrec.utils.sid.collision import concat_ranges -PROMPT_INPUT_IDS = "prompt_input_ids" -PROMPT_CU_SEQLENS = "prompt_cu_seqlens" -PROMPT_HOLE_POSITIONS = "prompt_hole_positions" -PROMPT_MAX_SEQLEN = "prompt_max_seqlen" -PROMPT_RESPONSE_LENGTHS = "prompt_response_lengths" +INPUT_IDS = "input_ids" +CU_SEQLENS = "cu_seqlens" +HOLE_POSITIONS = "hole_positions" +HOLE_KEYS = "hole_keys" +HOLE_SLOT_COUNTS = "hole_slot_counts" +RESPONSE_LENGTHS = "response_lengths" +MAX_SEQLEN = "max_seqlen" +# where the collator stores the streams on the batch +PROMPT_INFO_PREFIX = "prompt_" +PROMPT_INPUT_IDS = PROMPT_INFO_PREFIX + INPUT_IDS +PROMPT_CU_SEQLENS = PROMPT_INFO_PREFIX + CU_SEQLENS +PROMPT_HOLE_POSITIONS = PROMPT_INFO_PREFIX + HOLE_POSITIONS +PROMPT_MAX_SEQLEN = PROMPT_INFO_PREFIX + MAX_SEQLEN +PROMPT_RESPONSE_LENGTHS = PROMPT_INFO_PREFIX + RESPONSE_LENGTHS -@dataclass -class AssembledPrompt: - """One batch of assembled prompts, in packed varlen form. - Args: - input_ids: every sample's tokens concatenated, ``(total_tokens,)``. - cu_seqlens: sample boundaries into ``input_ids``, ``(batch_size + 1,)``. - hole_positions: absolute indices the projected embeddings overwrite, - grouped by included projected occurrence in - ``PromptPlan.projected_slots`` order, then by sample. - response_lengths: number of response tokens in each sample. - max_seqlen: widest sample, on the host so the model never derives it. +@torch.jit.script +def mix64(z: torch.Tensor) -> torch.Tensor: + """A SplitMix64-shaped avalanche over int64, wrapping. + + torch's right shift on a signed integer is arithmetic, so every shift the + mixer wants as logical is masked back. Getting that wrong is not a weaker + hash, it is a different function on negative inputs. """ + z = z * (-7046029254386353131) + z = (z ^ ((z >> 30) & 0x3FFFFFFFF)) * (-4658895280553007687) + z = (z ^ ((z >> 27) & 0x1FFFFFFFFF)) * (-7723592293110705685) + return z ^ ((z >> 31) & 0x1FFFFFFFF) + + +@torch.jit.script +def _exclusive_cumsum(values: torch.Tensor) -> torch.Tensor: + """Exclusive prefix sum along dim 0.""" + return torch.cumsum(values, dim=0) - values + + +@torch.jit.script +def _row_ids(lengths: torch.Tensor) -> torch.Tensor: + """Row index of every element in a jagged buffer described by lengths.""" + return torch.repeat_interleave( + torch.arange(lengths.numel(), dtype=torch.int64, device=lengths.device), + lengths, + ) - input_ids: np.ndarray - cu_seqlens: np.ndarray - hole_positions: np.ndarray - response_lengths: np.ndarray - max_seqlen: int +@torch.jit.script +def _within_row_index(lengths: torch.Tensor) -> torch.Tensor: + """Position of every element inside its own row.""" + total = int(torch.sum(lengths)) + starts = _exclusive_cumsum(lengths) + return torch.arange( + total, dtype=torch.int64, device=lengths.device + ) - torch.repeat_interleave(starts, lengths) -class PromptAssembler: - """Walks a ``PromptPlan`` to build token streams. +@torch.jit.script +def batch_device(batch: Dict[str, torch.Tensor]) -> torch.device: + """Device the batch already lives on, so nothing is built on the wrong one.""" + for value in batch.values(): + return value.device + return torch.device("cpu") + + +@torch.jit.script +def _destinations(seg_start: torch.Tensor, seg_len: torch.Tensor) -> torch.Tensor: + """Absolute index of every value of one segment in the packed stream.""" + return torch.repeat_interleave(seg_start, seg_len) + _within_row_index(seg_len) + + +def _wrap64(value: int) -> int: + """Reduce a Python int into the signed 64-bit range, wrapping. + + Every constant the fold mixes is built on the host and handed to torch as a + scalar, and torch refuses a scalar outside the tensor's dtype. Wrapping here + is what makes "the fold wraps" true at the boundary as well as inside it. + """ + value &= (1 << 64) - 1 + return value - (1 << 64) if value >= 1 << 63 else value + + +def _plan_salt(fold: FoldConstants, plan_hash: str) -> int: + """A 64-bit digest of ``plan_hash``, as a signed multiple of the plan constant.""" + if not plan_hash: + return 0 + digest = hashlib.sha256(plan_hash.encode("utf-8")).digest()[:8] + return _wrap64(int.from_bytes(digest, "little") * fold.plan) + + +def host_lengths_key(feature_name: str) -> str: + """The per-row item count a host lookup stage emits beside its rows.""" + return feature_name + "__lengths" + + +class PromptAssembler(nn.Module): + """Walks a compiled plan to build one batch's packed token stream. + + The collator calls it eagerly on the host; export scripts the same module + into the serving front-end. Validation runs at both call sites: an assembled + row is checked, never truncated or repaired, because a stream that is + silently wrong reaches the loss or the beam as plausible output. Args: - prompt_plan: the compiled walk order. - sid_space: resolved SID token space; required when a slot renders SIDs. - include_response: whether to read and emit the supervised response. + prompt_plan: the compiled walk order and its constants. + sid_space: resolved SID token space; required when a slot renders SID + codes or is projected. + plan_hash: the compiled plan's hash; its low bits salt every key, so + two plans cannot cross-match in a shared prefix cache. + include_response: whether to read and emit the supervised tail. """ + # TorchScript resolves a Final class attribute as a constant; a + # module-level one it cannot see at all + KIND_STATIC: Final[int] = 0 + KIND_INLINE: Final[int] = 1 + KIND_PROJECTED: Final[int] = 2 + # member index and position within a hole packed into one integer, wide + # enough that a position cannot carry into the member index + MEMBER_STRIDE: Final[int] = 1 << 32 + + kinds: List[int] + names: List[str] + static_tokens: List[List[int]] + exact_widths: List[int] + hole_slots: List[int] + is_sequences: List[bool] + salts: List[int] + member_names: List[List[str]] + level_lo: List[int] + level_hi: List[int] + def __init__( self, prompt_plan: PromptPlan, sid_space: Optional[ResolvedSidSpace] = None, + plan_hash: str = "", include_response: bool = True, ) -> None: - self._prompt_plan = prompt_plan - self._sid_space = sid_space - self._response_segments = ( - prompt_plan.response_segments if include_response else () - ) - if sid_space is not None: - self._flat_lo = np.asarray(sid_space.level_offsets, dtype=np.int64) - self._flat_hi = self._flat_lo + np.asarray( - sid_space.codebook, dtype=np.int64 - ) + super().__init__() + segments = tuple(prompt_plan.segments) + self.num_body = len(segments) + if include_response: + segments = segments + tuple(prompt_plan.response_segments) + inline = [ - s - for s in prompt_plan.segments + self._response_segments + s.name + for s in segments if isinstance(s, SlotSeg) and s.fill is FillMode.INLINE ] if inline and sid_space is None: raise ValueError( - f"prompt slots {[s.name for s in inline]} render INLINE, which " - f"means SID codes, but no sid_space was compiled." + f"prompt slots {inline} render INLINE, which means SID codes, but " + f"no sid_space was compiled." ) + self.sentinel = -1 + self.id_shift = 0 + self.num_levels = 1 + self.level_lo = [] + self.level_hi = [] + if sid_space is not None: + if sid_space.sentinel_token_id is not None: + self.sentinel = int(sid_space.sentinel_token_id) + self.id_shift = int(sid_space.base_vocab_size) + self.num_levels = int(sid_space.num_levels) + self.level_lo = [int(o) for o in sid_space.level_offsets] + self.level_hi = [ + int(o + c) for o, c in zip(sid_space.level_offsets, sid_space.codebook) + ] + self.max_length = int(prompt_plan.max_length) - def _inline_tokens( + self.kinds = [] + self.names = [] + self.static_tokens = [] + self.exact_widths = [] + self.hole_slots = [] + self.is_sequences = [] + self.member_names = [] + slot_ids: List[int] = [] + # holes are grouped by projected occurrence in emission order, which is + # the order of ``projected_slots`` and of the front-end's projections + occurrences = 0 + for index, seg in enumerate(segments): + if isinstance(seg, Static): + self._append( + self.KIND_STATIC, + "", + [int(t) for t in seg.token_ids], + -1, + -1, + False, + [], + ) + slot_ids.append(-1) + continue + assert isinstance(seg, SlotSeg) + is_sequence = seg.group_type == FeatureGroupType.JAGGED_SEQUENCE + if seg.fill is FillMode.INLINE: + # the answer's width sizes the loss window, so it is exact + width = -1 + if index >= self.num_body and seg.width.num_positions is not None: + width = int(seg.width.num_positions) + self._append( + self.KIND_INLINE, + seg.name, + [], + width, + -1, + is_sequence, + [seg.feature_names[0]], + ) + else: + if self.sentinel < 0: + raise ValueError( + f"prompt slot [{seg.name}] is PROJECTED but no sentinel token " + f"was compiled; a hole would be indistinguishable from content." + ) + self._append( + self.KIND_PROJECTED, + seg.name, + [], + -1, + occurrences, + is_sequence, + list(seg.feature_names), + ) + occurrences += 1 + slot_ids.append(int(seg.slot_id)) + + self.num_segments = len(self.kinds) + self.num_hole_slots = occurrences + fold = prompt_plan.fold + self.fold_value = int(fold.value) + self.fold_index = int(fold.index) + plan_salt = _plan_salt(fold, plan_hash) + self.salts = [_wrap64(fold.slot * slot_id + plan_salt) for slot_id in slot_ids] + + def _append( self, + kind: int, name: str, - flat: np.ndarray, - lengths: np.ndarray, - exact_width: Optional[int] = None, - ) -> np.ndarray: - """Validate offset SID codes against their bands and shift to token ids. - - The data carries ``level_offsets[l] + code``; the LM vocabulary needs - one further uniform shift by ``base_vocab_size``. + tokens: List[int], + exact_width: int, + hole_slot: int, + is_sequence: bool, + members: List[str], + ) -> None: + """Record one unrolled segment's constants.""" + self.kinds.append(kind) + self.names.append(name) + self.static_tokens.append(tokens) + self.exact_widths.append(exact_width) + self.hole_slots.append(hole_slot) + self.is_sequences.append(is_sequence) + self.member_names.append(members) - Args: - name: the slot, for the message. - flat: the slot's offset SID codes for the whole batch. - lengths: per-sample position counts into ``flat``. - exact_width: required position count, for the response; any whole - number of items otherwise. - """ - assert self._sid_space is not None - levels = self._sid_space.num_levels - partial = np.nonzero(lengths % levels)[0] - if partial.size: - sample = int(partial[0]) + def _lengths( + self, batch: Dict[str, torch.Tensor], slot: str, member: str, is_sequence: bool + ) -> torch.Tensor: + """Per-row item count, from the parsed dict or from a host lookup stage.""" + key = member + ".lengths" + if key in batch: + return batch[key].to(torch.int64) + host_key = host_lengths_key(member) + if host_key in batch: + return batch[host_key].to(torch.int64) + if is_sequence: raise ValueError( - f"prompt slot [{name}]: sample {sample} has " - f"{int(lengths[sample])} values, not a whole number of " - f"{levels}-level items." + "prompt slot [" + + slot + + "] reads [" + + member + + "], which the parser emitted as a scalar: a slot renders a " + + "sequence of SID codes, so the column must be list." ) - if exact_width is not None: - wrong = np.nonzero(lengths != exact_width)[0] - if wrong.size: - sample = int(wrong[0]) + # a dense member has one row per sample and no lengths + return torch.ones( + batch[member + ".values"].size(0), + dtype=torch.int64, + device=batch_device(batch), + ) + + def _batch_size(self, batch: Dict[str, torch.Tensor]) -> int: + """Row count, which every slot must agree on.""" + batch_size = -1 + for i in range(self.num_segments): + if self.kinds[i] == self.KIND_STATIC: + continue + for member in self.member_names[i]: + rows = int( + self._lengths( + batch, self.names[i], member, self.is_sequences[i] + ).numel() + ) + if batch_size < 0: + batch_size = rows + elif rows != batch_size: + raise ValueError( + "prompt slot [" + + self.names[i] + + "] has " + + str(rows) + + " samples, expected " + + str(batch_size) + + "." + ) + if batch_size >= 0: + return batch_size + if "batch_size" in batch: + return int(batch["batch_size"]) + return 0 + + def _inline_counts( + self, batch: Dict[str, torch.Tensor], index: int, batch_size: int + ) -> torch.Tensor: + """Token count each row contributes for one INLINE segment. + + For a multi-value sequence -- which is what a SID history is -- the row + holds ``lengths`` items and each item holds ``key_lengths`` codes, so the + count is a segmented sum rather than ``lengths`` itself. + """ + member = self.member_names[index][0] + lengths = self._lengths( + batch, self.names[index], member, self.is_sequences[index] + ) + key = member + ".key_lengths" + if key in batch: + key_lengths = batch[key].to(torch.int64).reshape(-1) + counts = torch.zeros(batch_size, dtype=torch.int64, device=lengths.device) + lengths = counts.index_add_(0, _row_ids(lengths), key_lengths) + return lengths + + def _inline_values( + self, batch: Dict[str, torch.Tensor], index: int, counts: torch.Tensor + ) -> torch.Tensor: + """Validate one INLINE segment's offset codes and shift them to token ids. + + The data carries ``level_offsets[l] + code``; the LM vocabulary needs one + further uniform shift by ``base_vocab_size``. + """ + name = self.names[index] + width = self.exact_widths[index] + if width >= 0: + wrong = torch.nonzero(counts != width) + if wrong.numel() > 0: + sample = int(wrong[0, 0]) raise ValueError( - f"prompt slot [{name}]: sample {sample} has " - f"{int(lengths[sample])} values, but the compiled width is " - f"{exact_width}. The loss window is sized from that width, " - f"so a wider row would be supervised only in part." + "prompt slot [" + + name + + "]: sample " + + str(sample) + + " has " + + str(int(counts[sample])) + + " values, but the compiled width is " + + str(width) + + ". The loss window is sized from that width, so a wider row " + + "would be supervised only in part." ) - by_level = flat.reshape(-1, levels) - if np.any(by_level < self._flat_lo) or np.any(by_level >= self._flat_hi): + partial = torch.nonzero(counts % self.num_levels != 0) + if partial.numel() > 0: + sample = int(partial[0, 0]) + raise ValueError( + "prompt slot [" + + name + + "]: sample " + + str(sample) + + " has " + + str(int(counts[sample])) + + " values, not a whole number of " + + str(self.num_levels) + + "-level items." + ) + values = ( + batch[self.member_names[index][0] + ".values"].to(torch.int64).reshape(-1) + ) + by_level = values.reshape(-1, self.num_levels) + lo = torch.tensor(self.level_lo, dtype=torch.int64, device=values.device) + hi = torch.tensor(self.level_hi, dtype=torch.int64, device=values.device) + if bool(torch.any(by_level < lo)) or bool(torch.any(by_level >= hi)): raise ValueError( - f"prompt slot [{name}]: SID values must already carry their " - f"level offset, so level l lies in " - f"[level_offsets[l], level_offsets[l] + codebook[l]). Read the " - f"offset_codebook column, not codebook or origin_codebook." + "prompt slot [" + + name + + "]: SID values must already carry their level offset, so level l " + + "lies in [level_offsets[l], level_offsets[l] + codebook[l]). Read " + + "the offset_codebook column, not codebook or origin_codebook." ) - return flat.astype(np.int64, copy=False) + self._sid_space.base_vocab_size + return values + self.id_shift - def _build_packed_prompt( + def _segment( self, - inline_flat: Dict[str, np.ndarray], - inline_lengths: Dict[str, np.ndarray], - projected_lengths: Dict[str, np.ndarray], + batch: Dict[str, torch.Tensor], + index: int, batch_size: int, - ) -> AssembledPrompt: - """Build the packed prompt in one pass over the compiled segments. + device: torch.device, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-row length and row-major values of one segment.""" + kind = self.kinds[index] + + if kind == self.KIND_STATIC: + run = torch.tensor( + self.static_tokens[index], dtype=torch.int64, device=device + ) + width = run.numel() + seg_len = torch.full((batch_size,), width, dtype=torch.int64, device=device) + return seg_len, run.unsqueeze(0).expand(batch_size, width).reshape(-1) + + if kind == self.KIND_INLINE: + counts = self._inline_counts(batch, index, batch_size) + return counts, self._inline_values(batch, index, counts) + + members = self.member_names[index] + name = self.names[index] + seg_len = self._lengths(batch, name, members[0], self.is_sequences[index]) + if self.is_sequences[index]: + for member in members[1:]: + other = self._lengths(batch, name, member, True) + if not torch.equal(seg_len, other): + raise ValueError( + "prompt slot [" + + name + + "] PROJECTED features [" + + members[0] + + "] and [" + + member + + "] have different per-sample lengths." + ) + else: + seg_len = torch.ones(batch_size, dtype=torch.int64, device=device) + total = int(torch.sum(seg_len)) + return seg_len, torch.full( + (total,), self.sentinel, dtype=torch.int64, device=seg_len.device + ) + + def _fold_segment( + self, + batch: Dict[str, torch.Tensor], + index: int, + batch_size: int, + hole_base: int, + keys: torch.Tensor, + ) -> None: + """Mix one projected segment's input values into ``keys``. + + Every member value that produces a hole contributes, discriminated by + slot, by member and by its index inside the hole. Without the last two + a two-member slot with values ``(a, b)`` would match one with + ``(b, a)``, and a permuted multi-value item would match itself + reordered -- both plausible, both wrong, and both silent. + """ + salt = self.salts[index] + name = self.names[index] + members = self.member_names[index] + is_sequence = self.is_sequences[index] + for member_index in range(len(members)): + member = members[member_index] + raw = batch[member + ".values"] + lengths = self._lengths(batch, name, member, is_sequence) + key_length_key = member + ".key_lengths" + + if not is_sequence: + # one hole per row; a dense member contributes its float32 bit + # pattern verbatim, which is the parsed input and not a + # computed reduction, so it is stable for a given request + if raw.is_floating_point(): + width = raw.size(1) + values = ( + raw.to(torch.float32) + .contiguous() + .view(torch.int32) + .to(torch.int64) + .reshape(-1) + & 0xFFFFFFFF + ) + hole = torch.repeat_interleave( + torch.arange(batch_size, dtype=torch.int64, device=raw.device), + torch.full( + (batch_size,), width, dtype=torch.int64, device=raw.device + ), + ) + local = ( + torch.arange(width, dtype=torch.int64, device=raw.device) + .unsqueeze(0) + .expand(batch_size, width) + .reshape(-1) + ) + else: + values = raw.to(torch.int64).reshape(-1) + hole = _row_ids(lengths) + local = _within_row_index(lengths) + else: + if raw.is_floating_point(): + raise ValueError( + "prompt slot [" + + name + + "] member [" + + member + + "] is a dense sequence; the fold has no per-item boundary " + + "for it." + ) + values = raw.to(torch.int64).reshape(-1) + if key_length_key in batch: + key_lengths = batch[key_length_key].to(torch.int64).reshape(-1) + hole = _row_ids(key_lengths) + local = _within_row_index(key_lengths) + else: + hole = torch.arange( + values.numel(), dtype=torch.int64, device=values.device + ) + local = torch.zeros_like(hole) + + local = local + member_index * self.MEMBER_STRIDE + mixed = mix64(values * self.fold_value + local * self.fold_index + salt) + keys.index_add_(0, hole + hole_base, mixed) - The loop is over segments, which the plan fixes, never over the batch: - each segment writes its whole column of the flat buffer at once. + def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + """Assemble one batch. Args: - inline_flat: INLINE slot name to its batch-wide value stream. - inline_lengths: INLINE slot name to its per-sample position counts. - projected_lengths: PROJECTED slot name to its per-sample counts. - batch_size: sample count. + batch: the parsed feature dict, keyed ``{feature}.values`` / + ``.lengths`` / ``.key_lengths`` as the data parser emits it. Returns: - The packed streams. + ``input_ids``, ``cu_seqlens``, ``hole_positions``, ``hole_keys``, + ``hole_slot_counts``, ``response_lengths`` and ``max_seqlen``. The + two hole-indexed streams are row-aligned: entry ``k`` of each + describes the same hole, grouped by projected occurrence in + emission order, then by sample. """ - segments = self._prompt_plan.segments + self._response_segments - body_count = len(self._prompt_plan.segments) + batch_size = self._batch_size(batch) + device = batch_device(batch) - seg_lengths = np.empty((len(segments), batch_size), dtype=np.int64) - for index, seg in enumerate(segments): - if isinstance(seg, Static): - seg_lengths[index] = len(seg.token_ids) - elif seg.fill is FillMode.INLINE: - seg_lengths[index] = inline_lengths[seg.name] - else: - seg_lengths[index] = projected_lengths[seg.name] - - row_lengths = seg_lengths.sum(axis=0) - cu_seqlens = np.concatenate(([0], np.cumsum(row_lengths))) - max_length = self._prompt_plan.max_length - if max_length: - over = np.nonzero(row_lengths > max_length)[0] - if over.size: - sample = int(over[0]) + seg_lens: List[torch.Tensor] = [] + seg_values: List[torch.Tensor] = [] + for i in range(self.num_segments): + length, values = self._segment(batch, i, batch_size, device) + seg_lens.append(length) + seg_values.append(values) + + stacked = torch.stack(seg_lens, dim=0) + row_total = torch.sum(stacked, dim=0) + if self.max_length > 0: + over = torch.nonzero(row_total > self.max_length) + if over.numel() > 0: + sample = int(over[0, 0]) raise ValueError( - f"assembled sample {sample} is {int(row_lengths[sample])} " - f"tokens, over max_length {max_length}. Samples are never " - f"truncated: cap the source features instead." + "assembled sample " + + str(sample) + + " is " + + str(int(row_total[sample])) + + " tokens, over max_length " + + str(self.max_length) + + ". Samples are never truncated: cap the source features instead." ) + row_start = _exclusive_cumsum(row_total) + seg_offsets = torch.cumsum(stacked, dim=0) - stacked - seg_starts = cu_seqlens[:-1] + (np.cumsum(seg_lengths, axis=0) - seg_lengths) + total_tokens = int(torch.sum(row_total)) + out = torch.zeros(total_tokens, dtype=torch.int64, device=device) - input_ids = np.empty(int(cu_seqlens[-1]), dtype=np.int64) - holes: List[np.ndarray] = [] - for index, seg in enumerate(segments): - lengths = seg_lengths[index] - destinations = concat_ranges(seg_starts[index], lengths) - if isinstance(seg, Static): - input_ids[destinations] = np.tile( - np.asarray(seg.token_ids, dtype=np.int64), batch_size - ) - elif seg.fill is FillMode.INLINE: - input_ids[destinations] = self._inline_tokens( - seg.name, - inline_flat[seg.name], - lengths, - exact_width=( - seg.width.num_positions if index >= body_count else None - ), - ) - else: - assert self._sid_space is not None - input_ids[destinations] = self._sid_space.sentinel_token_id - holes.append(destinations) - - response_lengths = seg_lengths[body_count:].sum(axis=0) - return AssembledPrompt( - input_ids=input_ids, - cu_seqlens=cu_seqlens, - hole_positions=( - np.concatenate(holes) if holes else np.empty(0, dtype=np.int64) - ), - response_lengths=response_lengths, - max_seqlen=int(row_lengths.max(initial=0)), - ) + # one entry per projected occurrence, and an empty stream under + # Pattern I, where the concatenation below would otherwise have + # nothing to join + hole_parts: List[torch.Tensor] = [ + torch.zeros(0, dtype=torch.int64, device=device) + ] + for _ in range(self.num_hole_slots): + hole_parts.append(torch.zeros(0, dtype=torch.int64, device=device)) - def forward( - self, parsed_features: Dict[str, "np.ndarray"] - ) -> Dict[str, np.ndarray]: - """Reshape one parsed batch, assemble it, and key it for the batch. + response_lengths = torch.zeros(batch_size, dtype=torch.int64, device=device) + for i in range(self.num_segments): + dest = _destinations(row_start + seg_offsets[i], seg_lens[i]) + out.index_copy_(0, dest, seg_values[i]) + slot = self.hole_slots[i] + if slot >= 0: + hole_parts[slot + 1] = dest + if i >= self.num_body: + response_lengths = response_lengths + seg_lens[i] - Args: - parsed_features: ``{column}.values`` / ``{column}.lengths`` as the - data parser emits them, for features and label fields alike. + hole_positions = torch.cat(hole_parts, dim=0) + # how many of those holes each projected occurrence owns, so a host can + # cut the flat streams back into per-slot spans without the plan + slot_counts = torch.zeros(self.num_hole_slots, dtype=torch.int64, device=device) + for slot in range(self.num_hole_slots): + slot_counts[slot] = hole_parts[slot + 1].numel() + keys = torch.zeros(hole_positions.numel(), dtype=torch.int64, device=device) + hole_base = 0 + for slot in range(self.num_hole_slots): + for i in range(self.num_segments): + if self.hole_slots[i] == slot: + self._fold_segment(batch, i, batch_size, hole_base, keys) + hole_base = hole_base + hole_parts[slot + 1].numel() - Returns: - The five streams, keyed as ``additional_infos`` expects them. - """ - inline_flat: Dict[str, np.ndarray] = {} - inline_lengths: Dict[str, np.ndarray] = {} - projected_lengths: Dict[str, np.ndarray] = {} - batch_size: Optional[int] = None - for seg in self._prompt_plan.segments + self._response_segments: - if not isinstance(seg, SlotSeg): - continue - sources = ( - seg.feature_names - if seg.fill is FillMode.PROJECTED - else seg.feature_names[:1] - ) - member_lengths: List[tuple[str, np.ndarray]] = [] - for source in sources: - lengths_key = f"{source}.lengths" - if seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: - if lengths_key not in parsed_features: - raise ValueError( - f"prompt slot [{seg.name}] reads [{source}], which " - f"the parser emitted as a scalar: a slot renders a " - f"sequence of SID codes, so the column must be " - f"list." - ) - lengths = np.asarray(parsed_features[lengths_key]) - slot_batch_size = int(lengths.size) - member_lengths.append((source, lengths)) - elif lengths_key in parsed_features: - slot_batch_size = int(np.asarray(parsed_features[lengths_key]).size) - else: - values = np.asarray(parsed_features[f"{source}.values"]) - slot_batch_size = int(values.shape[0]) - if batch_size is None: - batch_size = slot_batch_size - elif slot_batch_size != batch_size: - raise ValueError( - f"prompt slot [{seg.name}] has {slot_batch_size} samples, " - f"expected {batch_size}." - ) - if seg.fill is FillMode.INLINE: - source, lengths = member_lengths[0] - lengths = np.asarray(lengths, dtype=np.int64) - inline_flat[seg.name] = np.asarray( - parsed_features[f"{source}.values"] - ).reshape(-1) - # a multi-value sequence holds `lengths` items per sample and - # `key_lengths` codes per item, so a sample's position count is - # a segmented sum rather than its item count - key_lengths_key = f"{source}.key_lengths" - if key_lengths_key in parsed_features: - key_lengths = np.asarray( - parsed_features[key_lengths_key], dtype=np.int64 - ).reshape(-1) - per_sample = np.zeros(lengths.size, dtype=np.int64) - np.add.at( - per_sample, - np.repeat(np.arange(lengths.size), lengths), - key_lengths, - ) - lengths = per_sample - inline_lengths[seg.name] = lengths - elif seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: - source, lengths = member_lengths[0] - for other_source, other_lengths in member_lengths[1:]: - if not np.array_equal(lengths, other_lengths): - raise ValueError( - f"prompt slot [{seg.name}] PROJECTED features " - f"[{source}] and [{other_source}] have different " - "per-sample lengths." - ) - projected_lengths[seg.name] = lengths - else: - assert batch_size is not None - projected_lengths[seg.name] = np.ones(batch_size, dtype=np.int64) - - out = self._build_packed_prompt( - inline_flat, - inline_lengths, - projected_lengths, - batch_size if batch_size is not None else 0, + cu_seqlens = torch.cat( + [ + torch.zeros(1, dtype=torch.int64, device=device), + torch.cumsum(row_total, dim=0), + ] ) + if batch_size > 0: + max_seqlen = torch.max(row_total) + else: + max_seqlen = torch.zeros((), dtype=torch.int64, device=device) + # literals rather than the module constants above: TorchScript cannot + # see a module-level global. ``assembler_test`` pins the two together. return { - PROMPT_INPUT_IDS: out.input_ids, - PROMPT_CU_SEQLENS: out.cu_seqlens, - PROMPT_HOLE_POSITIONS: out.hole_positions, - PROMPT_MAX_SEQLEN: np.asarray(out.max_seqlen, dtype=np.int64), - PROMPT_RESPONSE_LENGTHS: out.response_lengths, + "input_ids": out, + "cu_seqlens": cu_seqlens.to(torch.int32), + "hole_positions": hole_positions, + "hole_keys": keys, + "hole_slot_counts": slot_counts, + "response_lengths": response_lengths, + "max_seqlen": max_seqlen, } diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index c6475b24..b121e9d6 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -12,14 +12,20 @@ import unittest import numpy as np +import torch from parameterized import parameterized +from tzrec.prompt import assembler from tzrec.prompt.assembler import ( - PROMPT_CU_SEQLENS, - PROMPT_HOLE_POSITIONS, - PROMPT_INPUT_IDS, - PROMPT_RESPONSE_LENGTHS, + CU_SEQLENS, + HOLE_KEYS, + HOLE_POSITIONS, + HOLE_SLOT_COUNTS, + INPUT_IDS, + MAX_SEQLEN, + RESPONSE_LENGTHS, PromptAssembler, + mix64, ) from tzrec.prompt.types import ( FillMode, @@ -67,9 +73,10 @@ def _slot( width_n=None, feature_names=None, group_type=FeatureGroupType.JAGGED_SEQUENCE, + slot_id=0, ) -> SlotSeg: return SlotSeg( - slot_id=0, + slot_id=slot_id, name=name, feature_names=tuple(feature_names) if feature_names is not None else (name,), group_type=group_type, @@ -118,23 +125,27 @@ def _parsed(inline=None, projected=None) -> dict: """ out = {} for name, rows in (inline or {}).items(): - out[f"{name}.values"] = ( + out[f"{name}.values"] = torch.as_tensor( np.concatenate(rows) if rows else np.zeros(0, dtype=np.int64) ) - out[f"{name}.lengths"] = np.asarray([len(row) for row in rows]) + out[f"{name}.lengths"] = torch.tensor([len(row) for row in rows]) for name, lengths in (projected or {}).items(): - out[f"{name}.lengths"] = np.asarray(lengths) + out[f"{name}.lengths"] = torch.tensor(lengths) return out +def _tensors(raw): + return {key: torch.as_tensor(np.asarray(value)) for key, value in raw.items()} + + class PromptAssemblerTest(unittest.TestCase): def test_inline_sid_gets_the_base_vocab_shift(self) -> None: asm = _asm((Static((7, 8)), _slot("hist", FillMode.INLINE))) # offset codes for one item: level 0 -> 1, level 1 -> 4+2, level 2 -> 8+3 - out = asm.forward(_parsed({"hist": [np.array([1, 6, 11])]})) + out = asm(_parsed({"hist": [np.array([1, 6, 11])]})) self.assertEqual( - out[PROMPT_INPUT_IDS].tolist(), + out[INPUT_IDS].tolist(), [ 7, 8, @@ -143,21 +154,30 @@ def test_inline_sid_gets_the_base_vocab_shift(self) -> None: _BASE_VOCAB_SIZE + 11, ], ) - self.assertEqual(out[PROMPT_CU_SEQLENS].tolist(), [0, 5]) - self.assertEqual(out[PROMPT_HOLE_POSITIONS].size, 0) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 5]) + self.assertEqual(out[HOLE_POSITIONS].numel(), 0) + self.assertEqual(int(out[MAX_SEQLEN]), 5) def test_projected_emits_sentinels_and_records_holes(self) -> None: asm = _asm((Static((7,)), _slot("prof", FillMode.PROJECTED, 4))) - out = asm.forward(_parsed(projected={"prof": [2, 3]})) + out = asm( + _parsed(projected={"prof": [2, 3]}) + | {"prof.values": torch.tensor([1, 2, 3, 4, 5])} + ) # sample 0: [7, S, S] sample 1: [7, S, S, S] self.assertEqual( - out[PROMPT_INPUT_IDS].tolist(), + out[INPUT_IDS].tolist(), [7, _SENTINEL, _SENTINEL, 7, _SENTINEL, _SENTINEL, _SENTINEL], ) - self.assertEqual(out[PROMPT_CU_SEQLENS].tolist(), [0, 3, 7]) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 3, 7]) # absolute indices into the flat buffer, which is what index_copy needs - self.assertEqual(out[PROMPT_HOLE_POSITIONS].tolist(), [1, 2, 4, 5, 6]) + self.assertEqual(out[HOLE_POSITIONS].tolist(), [1, 2, 4, 5, 6]) + self.assertEqual(out[HOLE_SLOT_COUNTS].tolist(), [5]) + holes = out[HOLE_POSITIONS] + self.assertTrue(bool(torch.all(out[INPUT_IDS][holes] == _SENTINEL))) + self.assertEqual(int(torch.sum(out[INPUT_IDS] == _SENTINEL)), 5) + self.assertEqual(int(out[MAX_SEQLEN]), 4) def test_holes_are_grouped_by_projected_occurrence_then_sample(self) -> None: plan = _plan( @@ -169,13 +189,15 @@ def test_holes_are_grouped_by_projected_occurrence_then_sample(self) -> None: ) ) asm = PromptAssembler(plan, _sid_space()) - out = asm.forward(_parsed(projected={"a": [1, 2], "b": [2, 1]})) - - self.assertEqual(out[PROMPT_CU_SEQLENS].tolist(), [0, 5, 11]) - self.assertEqual( - out[PROMPT_HOLE_POSITIONS].tolist(), [0, 5, 6, 2, 3, 8, 4, 9, 10] + out = asm( + _parsed(projected={"a": [1, 2], "b": [2, 1]}) + | {"a.values": torch.tensor([1, 2, 3]), "b.values": torch.tensor([4, 5, 6])} ) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 5, 11]) + self.assertEqual(out[HOLE_POSITIONS].tolist(), [0, 5, 6, 2, 3, 8, 4, 9, 10]) + self.assertEqual(out[HOLE_SLOT_COUNTS].tolist(), [3, 3, 3]) + def test_response_is_optional_and_its_length_is_recorded(self) -> None: plan = _plan( (Static((7,)), _slot("hist", FillMode.INLINE)), @@ -184,10 +206,10 @@ def test_response_is_optional_and_its_length_is_recorded(self) -> None: parsed = _parsed( {"hist": [np.array([1, 6, 11])], "answer": [np.array([0, 4, 8])]} ) - out = PromptAssembler(plan, _sid_space()).forward(parsed) + out = PromptAssembler(plan, _sid_space())(parsed) self.assertEqual( - out[PROMPT_INPUT_IDS].tolist(), + out[INPUT_IDS].tolist(), [ 7, _BASE_VOCAB_SIZE + 1, @@ -199,22 +221,54 @@ def test_response_is_optional_and_its_length_is_recorded(self) -> None: _BASE_VOCAB_SIZE + 8, ], ) - self.assertEqual(out[PROMPT_RESPONSE_LENGTHS].tolist(), [4]) + self.assertEqual(out[RESPONSE_LENGTHS].tolist(), [4]) - prompt_only = PromptAssembler( - plan, _sid_space(), include_response=False - ).forward(_parsed({"hist": [np.array([1, 6, 11])]})) + prompt_only = PromptAssembler(plan, _sid_space(), include_response=False)( + _parsed({"hist": [np.array([1, 6, 11])]}) + ) self.assertEqual( - prompt_only[PROMPT_INPUT_IDS].tolist(), + prompt_only[INPUT_IDS].tolist(), [7, _BASE_VOCAB_SIZE + 1, _BASE_VOCAB_SIZE + 6, _BASE_VOCAB_SIZE + 11], ) - self.assertEqual(prompt_only[PROMPT_RESPONSE_LENGTHS].tolist(), [0]) + self.assertEqual(prompt_only[RESPONSE_LENGTHS].tolist(), [0]) + + def test_rejects_a_response_of_the_wrong_width(self) -> None: + plan = _plan( + (_slot("hist", FillMode.INLINE),), + response=(_slot("answer", FillMode.INLINE),), + ) + asm = PromptAssembler(plan, _sid_space()) + with self.assertRaisesRegex(ValueError, "compiled width is 3"): + asm( + _parsed( + { + "hist": [np.array([1, 6, 11])], + "answer": [np.array([0, 4, 8, 1, 5, 9])], + } + ) + ) + + def test_a_multi_value_history_walks_like_the_flat_layout(self) -> None: + """Items with key_lengths and one code per position are one stream.""" + asm = _asm((Static((7,)), _slot("hist", FillMode.INLINE))) + flat = asm( + _parsed({"hist": [np.array([1, 6, 11, 2, 7, 10]), np.array([0, 4, 8])]}) + ) + items = asm( + { + "hist.values": torch.tensor([1, 6, 11, 2, 7, 10, 0, 4, 8]), + "hist.lengths": torch.tensor([2, 1]), + "hist.key_lengths": torch.tensor([3, 3, 3]), + } + ) + for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, MAX_SEQLEN): + self.assertTrue(torch.equal(flat[key], items[key]), key) def test_scalar_label_is_named_not_a_key_error(self) -> None: asm = _asm((_slot("hist", FillMode.INLINE),)) with self.assertRaisesRegex(ValueError, "must be\\s+list"): - asm.forward({"hist.values": np.array([1, 6, 11])}) + asm({"hist.values": torch.tensor([1, 6, 11])}) @parameterized.expand( [ @@ -226,17 +280,17 @@ def test_scalar_label_is_named_not_a_key_error(self) -> None: def test_rejects_a_code_outside_its_band(self, values) -> None: asm = _asm((_slot("hist", FillMode.INLINE),)) with self.assertRaisesRegex(ValueError, "offset_codebook column"): - asm.forward(_parsed({"hist": [np.array(values)]})) + asm(_parsed({"hist": [np.array(values)]})) def test_rejects_a_partial_item(self) -> None: asm = _asm((_slot("hist", FillMode.INLINE),)) with self.assertRaisesRegex(ValueError, "whole number of 3-level items"): - asm.forward(_parsed({"hist": [np.array([1, 6])]})) + asm(_parsed({"hist": [np.array([1, 6])]})) def test_over_long_row_is_an_error_not_a_truncation(self) -> None: asm = _asm((Static((7, 8, 9)), _slot("hist", FillMode.INLINE)), max_length=4) with self.assertRaisesRegex(ValueError, "never truncated"): - asm.forward(_parsed({"hist": [np.array([1, 6, 11])]})) + asm(_parsed({"hist": [np.array([1, 6, 11])]})) def test_inline_without_a_sid_space_is_rejected_at_construction(self) -> None: plan = _plan((_slot("hist", FillMode.INLINE),)) @@ -262,6 +316,7 @@ def test_column_shaped_values_are_flattened(self) -> None: out = assemble_into(compiled_prompt, parsed) self.assertEqual(out["prompt_cu_seqlens"].tolist(), [0, 3, 6]) self.assertEqual(out["prompt_input_ids"].tolist()[0], _BASE_VOCAB_SIZE + 1) + self.assertEqual(int(out["prompt_max_seqlen"]), 3) def test_rejects_inconsistent_slot_batch_sizes(self) -> None: plan = _plan( @@ -272,15 +327,15 @@ def test_rejects_inconsistent_slot_batch_sizes(self) -> None: ) asm = PromptAssembler(plan, _sid_space()) parsed = { - "hist.values": np.array([1, 6, 11]), - "hist.lengths": np.array([3]), - "answer.values": np.array([0, 4, 8, 1, 6, 11]), - "answer.lengths": np.array([3, 3]), + "hist.values": torch.tensor([1, 6, 11]), + "hist.lengths": torch.tensor([3]), + "answer.values": torch.tensor([0, 4, 8, 1, 6, 11]), + "answer.lengths": torch.tensor([3, 3]), } with self.assertRaisesRegex( ValueError, r"prompt slot \[answer\] has 2 samples, expected 1" ): - asm.forward(parsed) + asm(parsed) def test_rejects_mismatched_projected_member_lengths(self) -> None: plan = _plan( @@ -295,15 +350,15 @@ def test_rejects_mismatched_projected_member_lengths(self) -> None: ) asm = PromptAssembler(plan, _sid_space()) parsed = { - "age.lengths": np.array([2, 1]), - "country.lengths": np.array([2, 2]), + "age.lengths": torch.tensor([2, 1]), + "country.lengths": torch.tensor([2, 2]), } with self.assertRaisesRegex( ValueError, r"PROJECTED features \[age\] and \[country\] have different", ): - asm.forward(parsed) + asm(parsed) def test_deep_projected_members_emit_one_hole_per_sample(self) -> None: plan = _plan( @@ -317,17 +372,203 @@ def test_deep_projected_members_emit_one_hole_per_sample(self) -> None: ) ) asm = PromptAssembler(plan, _sid_space()) - out = asm.forward( + out = asm( { - "dense.values": np.array([[1.0, 2.0], [3.0, 4.0]]), - "sparse.values": np.array([5, 6, 7]), - "sparse.lengths": np.array([2, 1]), + "dense.values": torch.tensor([[1.0, 2.0], [3.0, 4.0]]), + "sparse.values": torch.tensor([5, 6, 7]), + "sparse.lengths": torch.tensor([2, 1]), } ) - self.assertEqual(out["prompt_input_ids"].tolist(), [_SENTINEL, _SENTINEL]) - self.assertEqual(out["prompt_cu_seqlens"].tolist(), [0, 1, 2]) - self.assertEqual(out["prompt_hole_positions"].tolist(), [0, 1]) + self.assertEqual(out[INPUT_IDS].tolist(), [_SENTINEL, _SENTINEL]) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 1, 2]) + self.assertEqual(out[HOLE_POSITIONS].tolist(), [0, 1]) + + def test_output_keys_match_the_module_constants(self) -> None: + """``forward`` writes literals; they must equal the exported names.""" + module = PromptAssembler(_plan((Static((7,)),))) + out = module({"batch_size": torch.tensor(2)}) + self.assertEqual( + set(out.keys()), + { + INPUT_IDS, + CU_SEQLENS, + HOLE_POSITIONS, + HOLE_KEYS, + HOLE_SLOT_COUNTS, + RESPONSE_LENGTHS, + MAX_SEQLEN, + }, + ) + self.assertEqual(out[INPUT_IDS].tolist(), [7, 7]) + self.assertEqual(assembler.PROMPT_INPUT_IDS, "prompt_" + INPUT_IDS) + + def test_scripting_preserves_every_output(self) -> None: + """The artifact and the collator's module are the same function.""" + plan = _plan( + ( + Static((7,)), + _slot("hist", FillMode.INLINE), + _slot("beh", FillMode.PROJECTED, 4), + ), + response=(_slot("answer", FillMode.INLINE),), + ) + module = PromptAssembler(plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") + batch = { + "hist.values": torch.tensor([1, 6, 11, 2, 7, 10]), + "hist.lengths": torch.tensor([1, 1]), + "hist.key_lengths": torch.tensor([3, 3]), + "beh.values": torch.tensor([7, 8, 9, 21, 22]), + "beh.lengths": torch.tensor([3, 2]), + "answer.values": torch.tensor([3, 7, 11, 0, 4, 8]), + "answer.lengths": torch.tensor([3, 3]), + } + eager = module(batch) + scripted = torch.jit.script(module) + for key, value in scripted(batch).items(): + self.assertTrue(torch.equal(eager[key], value), key) + + def test_a_scripted_walk_reports_its_validation(self) -> None: + """The exported artifact refuses a bad request with the same message.""" + scripted = torch.jit.script( + _asm((Static((7, 8, 9)), _slot("hist", FillMode.INLINE)), max_length=4) + ) + with self.assertRaisesRegex(torch.jit.Error, "never truncated"): + scripted(_parsed({"hist": [np.array([1, 6, 11])]})) + + +class MixTest(unittest.TestCase): + def test_mix64_matches_a_host_reference(self) -> None: + """The mixer's masked shifts must reproduce SplitMix64 on negatives.""" + + def reference(value: int) -> int: + mask = (1 << 64) - 1 + z = (value * 0x9E3779B97F4A7C15) & mask + z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & mask + z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & mask + z = z ^ (z >> 31) + return z - (1 << 64) if z >= 1 << 63 else z + + values = [0, 1, -1, 2**31, -(2**31), 123456789, -987654321] + got = mix64(torch.tensor(values, dtype=torch.int64)).tolist() + self.assertEqual(got, [reference(v) for v in values]) + + +class FoldTest(unittest.TestCase): + def setUp(self) -> None: + self.beh = _slot("beh", FillMode.PROJECTED, 30) + self.plan = _plan((self.beh,)) + self.module = PromptAssembler( + self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718" + ) + + def _keys(self, values, lengths): + return self.module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array(lengths, dtype=np.int64), + } + ) + )[HOLE_KEYS] + + def test_equal_content_folds_equal(self) -> None: + """A key is a function of the hole's inputs and nothing else.""" + self.assertTrue( + torch.equal(self._keys([5, 6, 7], [3]), self._keys([5, 6, 7], [3])) + ) + + def test_different_content_folds_apart(self) -> None: + """The whole point: a changed input must not reuse the cached KV.""" + first = self._keys([5, 6, 7], [3]) + second = self._keys([5, 6, 8], [3]) + self.assertEqual(first[:2].tolist(), second[:2].tolist()) + self.assertNotEqual(int(first[2]), int(second[2])) + + def test_the_plan_hash_salts_the_keys(self) -> None: + """Two plans must not cross-match in a shared prefix cache.""" + other = PromptAssembler(self.plan, _sid_space(), plan_hash="ffffffffffffffff") + batch = _tensors( + { + "beh.values": np.array([5, 6, 7], dtype=np.int64), + "beh.lengths": np.array([3], dtype=np.int64), + } + ) + self.assertFalse( + torch.equal(self.module(batch)[HOLE_KEYS], other(batch)[HOLE_KEYS]) + ) + + def test_a_permuted_multi_value_item_does_not_collide(self) -> None: + """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" + + def keys(values): + return self.module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array([1], dtype=np.int64), + "beh.key_lengths": np.array([3], dtype=np.int64), + } + ) + )[HOLE_KEYS] + + self.assertNotEqual(keys([1, 5, 9]).tolist(), keys([9, 5, 1]).tolist()) + + def test_two_members_exchanging_values_do_not_collide(self) -> None: + """Without the member index a two-member slot is order-blind.""" + slot = _slot("pair", FillMode.PROJECTED, 30, feature_names=("a", "b")) + module = PromptAssembler( + _plan((slot,)), _sid_space(), plan_hash="a1b2c3d4e5f60718" + ) + + def keys(first, second): + return module( + _tensors( + { + "a.values": np.array(first, dtype=np.int64), + "a.lengths": np.array([1], dtype=np.int64), + "b.values": np.array(second, dtype=np.int64), + "b.lengths": np.array([1], dtype=np.int64), + } + ) + )[HOLE_KEYS] + + self.assertNotEqual(keys([3], [9]).tolist(), keys([9], [3]).tolist()) + + def test_a_dense_member_folds_its_bit_pattern(self) -> None: + """A float member contributes the parsed input verbatim, so it is stable.""" + slot = _slot("vec", FillMode.PROJECTED, group_type=FeatureGroupType.DEEP) + module = PromptAssembler(_plan((slot,)), _sid_space(), plan_hash="a1b2") + + def keys(rows): + return module({"vec.values": torch.tensor(rows, dtype=torch.float32)})[ + HOLE_KEYS + ] + + self.assertEqual(keys([[0.5, 1.0]]).tolist(), keys([[0.5, 1.0]]).tolist()) + self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) + + def test_a_dense_sequence_member_is_rejected(self) -> None: + with self.assertRaisesRegex(ValueError, "no per-item boundary"): + self.module( + { + "beh.values": torch.tensor([[0.5], [1.0]]), + "beh.lengths": torch.tensor([2]), + } + ) + + @unittest.skipIf(not torch.cuda.is_available(), "no GPU") + def test_the_fold_is_bit_identical_across_devices(self) -> None: + """Integer addition cannot depend on the order a device reduces in.""" + batch = _tensors( + { + "beh.values": np.arange(64, dtype=np.int64), + "beh.lengths": np.array([32, 32], dtype=np.int64), + } + ) + on_cpu = self.module(batch)[HOLE_KEYS] + on_gpu = self.module({k: v.cuda() for k, v in batch.items()})[HOLE_KEYS] + self.assertTrue(torch.equal(on_cpu, on_gpu.cpu())) if __name__ == "__main__": diff --git a/tzrec/prompt/compile.py b/tzrec/prompt/compile.py index 030e3272..e9244b5e 100644 --- a/tzrec/prompt/compile.py +++ b/tzrec/prompt/compile.py @@ -403,7 +403,6 @@ def compile_prompt( logits_suffix_len=_suffix_keep(response), static_prefix_len=_static_prefix_len(body), projected_slots=projected, - slot_index={seg.name: index for index, seg in enumerate(projected)}, ) _validate(plan) diff --git a/tzrec/prompt/export.py b/tzrec/prompt/export.py index 80b266da..2d4027b2 100644 --- a/tzrec/prompt/export.py +++ b/tzrec/prompt/export.py @@ -46,19 +46,18 @@ from tzrec.acc import utils as acc_utils from tzrec.features.feature import BaseFeature, create_feature_configs, create_fg_json -from tzrec.prompt.frontend import ( - OUT_CU_SEQLENS, - OUT_HOLE_KEYS, - OUT_HOLE_POSITIONS, - OUT_HOLE_SLOT_COUNTS, - OUT_INPUT_IDS, - OUT_RESPONSE_LENGTHS, - OUT_SLOT_EMBEDS, +from tzrec.prompt.assembler import ( + CU_SEQLENS, + HOLE_KEYS, + HOLE_POSITIONS, + HOLE_SLOT_COUNTS, + INPUT_IDS, + MAX_SEQLEN, + RESPONSE_LENGTHS, PromptAssembler, - PromptFrontEnd, - SlotTable, host_lengths_key, ) +from tzrec.prompt.frontend import SLOT_EMBEDS, PromptFrontEnd, SlotTable from tzrec.prompt.types import CompiledPrompt, FillMode, SlotSeg from tzrec.protos.model_pb2 import FeatureGroupType from tzrec.protos.pipeline_pb2 import EasyRecConfig @@ -294,8 +293,6 @@ def build_front_end( assembler = PromptAssembler( plan, sid_space, - features_are_dense={f.name: not f.is_sparse for f in features}, - features_are_multi_valued={f.name: _is_multi_valued(f) for f in features}, plan_hash=compiled_prompt.plan_hash, # serving has no answer to assemble; the LM generates it include_response=False, @@ -327,13 +324,14 @@ def build_front_end( "lookup": lookup, "inputs": _frontend_inputs(compiled_prompt, features, lookup, embed_keys), "outputs": [ - OUT_INPUT_IDS, - OUT_CU_SEQLENS, - OUT_HOLE_POSITIONS, - OUT_HOLE_KEYS, - OUT_HOLE_SLOT_COUNTS, - OUT_SLOT_EMBEDS, - OUT_RESPONSE_LENGTHS, + INPUT_IDS, + CU_SEQLENS, + HOLE_POSITIONS, + HOLE_KEYS, + HOLE_SLOT_COUNTS, + SLOT_EMBEDS, + RESPONSE_LENGTHS, + MAX_SEQLEN, ], } return front_end, meta, dense_meta diff --git a/tzrec/prompt/frontend.py b/tzrec/prompt/frontend.py index 338d476b..4d0852e9 100644 --- a/tzrec/prompt/frontend.py +++ b/tzrec/prompt/frontend.py @@ -11,544 +11,20 @@ """The prompt front-end a serving runtime loads: assemble, look up, project. -The assembler is the collator's walk as tensor ops. ``PromptPlan`` is a -compile-time constant, so the segment loop unrolls into parallel constant -lists at construction and what remains is jagged integer arithmetic -- -``cumsum``, ``repeat_interleave``, ``index_copy_`` -- with no data-dependent -control flow, which is what lets ``torch.jit.script`` carry it into a runtime -that has no tzrec source. - -``hole_keys`` is integer end to end: ``int64`` addition is associative and -commutative and wraps deterministically, so the fold cannot depend on the order -``index_add_`` happens to reduce in, on the device, or on how the batch was -split. A float accumulator would satisfy "do not fold the projected vector" in -letter and reintroduce the variance in spirit. +It wraps the collator's own ``PromptAssembler`` -- the one walk, scripted -- +and adds the two stages serving needs after it: the slot lookup, when the host +has no embedding stage of its own, and the trained projections into the LM +input space. """ -from typing import Dict, Final, List, Optional, Tuple +from typing import Dict, List, Optional import torch from torch import nn -from tzrec.prompt.types import ( - FillMode, - FoldConstants, - PromptPlan, - ResolvedSidSpace, - SlotSeg, - Static, -) -from tzrec.protos.model_pb2 import FeatureGroupType - -OUT_INPUT_IDS = "input_ids" -OUT_CU_SEQLENS = "cu_seqlens" -OUT_HOLE_POSITIONS = "hole_positions" -OUT_HOLE_KEYS = "hole_keys" -OUT_HOLE_SLOT_COUNTS = "hole_slot_counts" -OUT_SLOT_EMBEDS = "slot_embeds" -OUT_RESPONSE_LENGTHS = "response_lengths" - - -@torch.jit.script -def mix64(z: torch.Tensor) -> torch.Tensor: - """A SplitMix64-shaped avalanche over int64, wrapping. - - torch's right shift on a signed integer is arithmetic, so every shift the - mixer wants as logical is masked back. Getting that wrong is not a weaker - hash, it is a different function on negative inputs. - """ - z = z * (-7046029254386353131) - z = (z ^ ((z >> 30) & 0x3FFFFFFFF)) * (-4658895280553007687) - z = (z ^ ((z >> 27) & 0x1FFFFFFFFF)) * (-7723592293110705685) - return z ^ ((z >> 31) & 0x1FFFFFFFF) - - -@torch.jit.script -def _exclusive_cumsum(values: torch.Tensor) -> torch.Tensor: - """Exclusive prefix sum along dim 0.""" - return torch.cumsum(values, dim=0) - values - - -@torch.jit.script -def _row_ids(lengths: torch.Tensor) -> torch.Tensor: - """Row index of every element in a jagged buffer described by lengths.""" - return torch.repeat_interleave( - torch.arange(lengths.numel(), dtype=torch.int64, device=lengths.device), - lengths, - ) - - -@torch.jit.script -def _within_row_index(lengths: torch.Tensor) -> torch.Tensor: - """Position of every element inside its own row.""" - total = int(torch.sum(lengths)) - starts = _exclusive_cumsum(lengths) - return torch.arange( - total, dtype=torch.int64, device=lengths.device - ) - torch.repeat_interleave(starts, lengths) - - -@torch.jit.script -def _pick_device(batch: Dict[str, torch.Tensor]) -> torch.device: - """Device the batch already lives on, so nothing is built on the wrong one.""" - for value in batch.values(): - return value.device - return torch.device("cpu") - - -@torch.jit.script -def _destinations(seg_start: torch.Tensor, seg_len: torch.Tensor) -> torch.Tensor: - """Absolute index of every value of one segment in the packed stream.""" - return torch.repeat_interleave(seg_start, seg_len) + _within_row_index(seg_len) - - -def _wrap64(value: int) -> int: - """Reduce a Python int into the signed 64-bit range, wrapping. - - Every constant the fold mixes is built on the host and handed to torch as a - scalar, and torch refuses a scalar outside the tensor's dtype. Wrapping here - is what makes "the fold wraps" true at the boundary as well as inside it. - """ - value &= (1 << 64) - 1 - return value - (1 << 64) if value >= 1 << 63 else value - - -def _plan_salt(fold: FoldConstants, plan_hash: str) -> int: - """Low 64 bits of ``plan_hash``, as a signed multiple of the plan constant.""" - if not plan_hash: - return 0 - return _wrap64(_wrap64(int(plan_hash[:16], 16)) * fold.plan) - - -def host_lengths_key(feature_name: str) -> str: - """The per-row item count a host lookup stage emits beside its rows.""" - return f"{feature_name}__lengths" - - -class PromptAssembler(nn.Module): - """Walks a compiled plan to build one batch's packed token stream. - - Args: - plan: the compiled walk order and its constants. - sid_space: resolved SID token space; required when a slot renders SID - codes or is projected. - features_are_dense: feature name to whether it arrives as floats. - features_are_multi_valued: feature name to whether it carries - ``key_lengths``, that is whether ``value_dim != 1``. - plan_hash: the compiled plan's hash; its low bits salt every key, so - two plans cannot cross-match in a shared prefix cache. - include_response: whether to emit the supervised tail. - """ - - # TorchScript resolves a Final class attribute as a constant; a - # module-level one it cannot see at all - KIND_STATIC: Final[int] = 0 - KIND_INLINE: Final[int] = 1 - KIND_PROJECTED: Final[int] = 2 - # member index and position within a hole packed into one integer, wide - # enough that a position cannot carry into the member index - MEMBER_STRIDE: Final[int] = 1 << 32 - - kinds: List[int] - static_tokens: List[List[int]] - value_keys: List[str] - length_keys: List[str] - host_length_keys: List[str] - key_length_keys: List[str] - hole_slots: List[int] - slot_ids: List[int] - is_sequences: List[bool] - salts: List[int] - member_value_keys: List[List[str]] - member_length_keys: List[List[str]] - member_host_length_keys: List[List[str]] - member_key_length_keys: List[List[str]] - member_is_dense: List[bool] - - def __init__( - self, - plan: PromptPlan, - sid_space: Optional[ResolvedSidSpace] = None, - features_are_dense: Optional[Dict[str, bool]] = None, - features_are_multi_valued: Optional[Dict[str, bool]] = None, - plan_hash: str = "", - include_response: bool = True, - ) -> None: - super().__init__() - dense = features_are_dense if features_are_dense is not None else {} - multi = ( - features_are_multi_valued if features_are_multi_valued is not None else {} - ) - segments = tuple(plan.segments) - self.num_body = len(segments) - if include_response: - segments = segments + tuple(plan.response_segments) - - sentinel = -1 - if sid_space is not None and sid_space.sentinel_token_id is not None: - sentinel = int(sid_space.sentinel_token_id) - self.sentinel = sentinel - id_shift = 0 if sid_space is None else int(sid_space.base_vocab_size) - - self.kinds = [] - self.static_tokens = [] - self.value_keys = [] - self.length_keys = [] - self.host_length_keys = [] - self.key_length_keys = [] - self.hole_slots = [] - self.slot_ids = [] - self.is_sequences = [] - self.member_value_keys = [] - self.member_length_keys = [] - self.member_host_length_keys = [] - self.member_key_length_keys = [] - self.member_is_dense = [] - - for seg in segments: - if isinstance(seg, Static): - self._append( - PromptAssembler.KIND_STATIC, - [int(t) for t in seg.token_ids], - "", - "", - "", - "", - -1, - -1, - False, - [], - [], - [], - [], - ) - continue - - assert isinstance(seg, SlotSeg) - primary = seg.feature_names[0] - is_sequence = seg.group_type == FeatureGroupType.JAGGED_SEQUENCE - if seg.fill is FillMode.INLINE: - if sid_space is None: - raise ValueError( - f"prompt slot [{seg.name}] renders INLINE, which means " - f"SID codes, but no sid_space was compiled." - ) - self._append( - PromptAssembler.KIND_INLINE, - [], - f"{primary}.values", - f"{primary}.lengths", - host_lengths_key(primary), - f"{primary}.key_lengths" if multi.get(primary, False) else "", - -1, - int(seg.slot_id), - is_sequence, - [], - [], - [], - [], - ) - continue - - if sentinel < 0: - raise ValueError( - f"prompt slot [{seg.name}] is PROJECTED but no sentinel token " - f"was compiled; a hole would be indistinguishable from content." - ) - for name in seg.feature_names: - if dense.get(name, False) and is_sequence: - raise ValueError( - f"prompt slot [{seg.name}] member [{name}] is a dense " - f"sequence; the fold has no per-item boundary for it." - ) - self._append( - PromptAssembler.KIND_PROJECTED, - [], - "", - f"{primary}.lengths", - host_lengths_key(primary), - "", - int(plan.slot_index[seg.name]), - int(seg.slot_id), - is_sequence, - [f"{name}.values" for name in seg.feature_names], - [f"{name}.lengths" for name in seg.feature_names], - [host_lengths_key(name) for name in seg.feature_names], - [ - f"{name}.key_lengths" if multi.get(name, False) else "" - for name in seg.feature_names - ], - ) - self.member_is_dense = self.member_is_dense + [ - dense.get(name, False) for name in seg.feature_names - ] - - self.id_shift = id_shift - self.num_segments = len(self.kinds) - self.num_hole_slots = len(plan.projected_slots) - self.fold_value = int(plan.fold.value) - self.fold_index = int(plan.fold.index) - plan_salt = _plan_salt(plan.fold, plan_hash) - self.salts = [ - _wrap64(plan.fold.slot * slot_id + plan_salt) for slot_id in self.slot_ids - ] - - def _append( - self, - kind: int, - tokens: List[int], - value_key: str, - length_key: str, - host_length_key: str, - key_length_key: str, - hole_slot: int, - slot_id: int, - is_sequence: bool, - member_values: List[str], - member_lengths: List[str], - member_host_lengths: List[str], - member_key_lengths: List[str], - ) -> None: - """Record one unrolled segment's constants.""" - self.kinds.append(kind) - self.static_tokens.append(tokens) - self.value_keys.append(value_key) - self.length_keys.append(length_key) - self.host_length_keys.append(host_length_key) - self.key_length_keys.append(key_length_key) - self.hole_slots.append(hole_slot) - self.slot_ids.append(slot_id) - self.is_sequences.append(is_sequence) - self.member_value_keys.append(member_values) - self.member_length_keys.append(member_lengths) - self.member_host_length_keys.append(member_host_lengths) - self.member_key_length_keys.append(member_key_lengths) - - def _lengths( - self, batch: Dict[str, torch.Tensor], key: str, host_key: str - ) -> torch.Tensor: - """Per-row item count, from the parsed dict or from a host lookup stage.""" - if key in batch: - return batch[key].to(torch.int64) - return batch[host_key].to(torch.int64) - - def _batch_size(self, batch: Dict[str, torch.Tensor]) -> int: - """Row count, from the first segment that names a per-row length.""" - for i in range(self.num_segments): - key = self.length_keys[i] - if key != "": - if key in batch: - return int(batch[key].numel()) - host_key = self.host_length_keys[i] - if host_key in batch: - return int(batch[host_key].numel()) - if "batch_size" in batch: - return int(batch["batch_size"]) - return 0 - - def _values_per_row( - self, batch: Dict[str, torch.Tensor], index: int, batch_size: int - ) -> torch.Tensor: - """Token count each row contributes for one INLINE segment. - - For a multi-value sequence -- which is what a SID history is -- the row - holds ``lengths`` items and each item holds ``key_lengths`` values, so - the count is a segmented sum rather than ``lengths`` itself. - """ - lengths = self._lengths( - batch, self.length_keys[index], self.host_length_keys[index] - ) - key = self.key_length_keys[index] - if key == "": - return lengths - key_lengths = batch[key].to(torch.int64) - counts = torch.zeros(batch_size, dtype=torch.int64, device=lengths.device) - return counts.index_add_(0, _row_ids(lengths), key_lengths) - - def _segment( - self, - batch: Dict[str, torch.Tensor], - index: int, - batch_size: int, - device: torch.device, - ) -> Tuple[torch.Tensor, torch.Tensor]: - """Per-row length and row-major values of one segment.""" - kind = self.kinds[index] - - if kind == self.KIND_STATIC: - run = torch.tensor( - self.static_tokens[index], dtype=torch.int64, device=device - ) - width = run.numel() - seg_len = torch.full((batch_size,), width, dtype=torch.int64, device=device) - return seg_len, run.unsqueeze(0).expand(batch_size, width).reshape(-1) - - if kind == self.KIND_INLINE: - seg_len = self._values_per_row(batch, index, batch_size) - values = batch[self.value_keys[index]].to(torch.int64).reshape(-1) - return seg_len, values + self.id_shift - - if self.is_sequences[index]: - seg_len = self._lengths( - batch, self.length_keys[index], self.host_length_keys[index] - ) - else: - seg_len = torch.ones(batch_size, dtype=torch.int64, device=device) - total = int(torch.sum(seg_len)) - return seg_len, torch.full( - (total,), self.sentinel, dtype=torch.int64, device=seg_len.device - ) - - def _fold_segment( - self, - batch: Dict[str, torch.Tensor], - index: int, - batch_size: int, - hole_base: int, - keys: torch.Tensor, - ) -> None: - """Mix one projected segment's input values into ``keys``. - - Every member value that produces a hole contributes, discriminated by - slot, by member and by its index inside the hole. Without the last two - a two-member slot with values ``(a, b)`` would match one with - ``(b, a)``, and a permuted multi-value item would match itself - reordered -- both plausible, both wrong, and both silent. - """ - salt = self.salts[index] - member_values = self.member_value_keys[index] - for member in range(len(member_values)): - raw = batch[member_values[member]] - lengths = self._lengths( - batch, - self.member_length_keys[index][member], - self.member_host_length_keys[index][member], - ) - key_length_key = self.member_key_length_keys[index][member] - - if not self.is_sequences[index]: - # one hole per row; a dense member contributes its float32 bit - # pattern verbatim, which is the parsed input and not a - # computed reduction, so it is stable for a given request - if raw.dtype == torch.float32: - width = raw.size(1) - values = ( - raw.contiguous().view(torch.int32).to(torch.int64).reshape(-1) - & 0xFFFFFFFF - ) - hole = torch.repeat_interleave( - torch.arange(batch_size, dtype=torch.int64, device=raw.device), - torch.full( - (batch_size,), width, dtype=torch.int64, device=raw.device - ), - ) - local = ( - torch.arange(width, dtype=torch.int64, device=raw.device) - .unsqueeze(0) - .expand(batch_size, width) - .reshape(-1) - ) - else: - values = raw.to(torch.int64).reshape(-1) - hole = _row_ids(lengths) - local = _within_row_index(lengths) - else: - values = raw.to(torch.int64).reshape(-1) - if key_length_key == "": - hole = torch.arange( - values.numel(), dtype=torch.int64, device=values.device - ) - local = torch.zeros_like(hole) - else: - key_lengths = batch[key_length_key].to(torch.int64) - hole = _row_ids(key_lengths) - local = _within_row_index(key_lengths) - - local = local + member * self.MEMBER_STRIDE - mixed = mix64(values * self.fold_value + local * self.fold_index + salt) - keys.index_add_(0, hole + hole_base, mixed) - - def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: - """Assemble one batch. - - Args: - batch: the parsed feature dict, keyed ``{feature}.values`` / - ``.lengths`` / ``.key_lengths`` as both hosts already emit it. - - Returns: - ``input_ids``, ``cu_seqlens``, ``hole_positions``, ``hole_keys``, - ``hole_slot_counts`` and ``response_lengths``. The two hole-indexed - outputs are row-aligned: entry ``k`` of each describes the same - hole, in ``projected_slots`` order. - """ - batch_size = self._batch_size(batch) - device = _pick_device(batch) - - seg_lens: List[torch.Tensor] = [] - seg_values: List[torch.Tensor] = [] - for i in range(self.num_segments): - length, values = self._segment(batch, i, batch_size, device) - seg_lens.append(length) - seg_values.append(values) - - stacked = torch.stack(seg_lens, dim=0) - row_total = torch.sum(stacked, dim=0) - row_start = _exclusive_cumsum(row_total) - seg_offsets = torch.cumsum(stacked, dim=0) - stacked - - total_tokens = int(torch.sum(row_total)) - out = torch.zeros(total_tokens, dtype=torch.int64, device=row_total.device) - - # one entry per projected slot, and an empty stream under Pattern I, - # where the concatenation below would otherwise have nothing to join - hole_parts: List[torch.Tensor] = [ - torch.zeros(0, dtype=torch.int64, device=device) - ] - for _ in range(self.num_hole_slots): - hole_parts.append(torch.zeros(0, dtype=torch.int64, device=device)) +from tzrec.prompt.assembler import PromptAssembler, batch_device - response_lengths = torch.zeros( - batch_size, dtype=torch.int64, device=row_total.device - ) - for i in range(self.num_segments): - dest = _destinations(row_start + seg_offsets[i], seg_lens[i]) - out.index_copy_(0, dest, seg_values[i]) - slot = self.hole_slots[i] - if slot >= 0: - hole_parts[slot + 1] = dest - if i >= self.num_body: - response_lengths = response_lengths + seg_lens[i] - - hole_positions = torch.cat(hole_parts, dim=0) - # how many of those holes each projected slot owns, so a host can cut - # the flat streams back into per-slot spans without re-deriving the plan - slot_counts = torch.zeros(self.num_hole_slots, dtype=torch.int64, device=device) - for slot in range(self.num_hole_slots): - slot_counts[slot] = hole_parts[slot + 1].numel() - keys = torch.zeros(hole_positions.numel(), dtype=torch.int64, device=out.device) - hole_base = 0 - for slot in range(self.num_hole_slots): - for i in range(self.num_segments): - if self.hole_slots[i] == slot: - self._fold_segment(batch, i, batch_size, hole_base, keys) - hole_base = hole_base + hole_parts[slot + 1].numel() - - cu_seqlens = torch.cat( - [ - torch.zeros(1, dtype=torch.int64, device=row_total.device), - torch.cumsum(row_total, dim=0), - ] - ) - # literals rather than the module constants above: TorchScript cannot - # see a module-level global. ``frontend_test`` pins the two together. - return { - "input_ids": out, - "cu_seqlens": cu_seqlens.to(torch.int32), - "hole_positions": hole_positions, - "hole_keys": keys, - "hole_slot_counts": slot_counts, - "response_lengths": response_lengths, - } +SLOT_EMBEDS = "slot_embeds" class SlotTable(nn.Module): @@ -703,7 +179,7 @@ def forward( Returns: The assembler's outputs plus ``slot_embeds``. """ - target = _pick_device(data) if device is None else device + target = batch_device(data) if device is None else device batch: Dict[str, torch.Tensor] = {} for key, value in data.items(): batch[key] = value.to(target) @@ -728,6 +204,8 @@ def forward( parts.append(projection(features[index])) index += 1 + # literals rather than the module constants: TorchScript cannot see a + # module-level global. ``frontend_test`` pins them together. if len(parts) > 0: out["slot_embeds"] = torch.cat(parts, dim=0) else: diff --git a/tzrec/prompt/frontend_test.py b/tzrec/prompt/frontend_test.py index f771cc42..a171a113 100644 --- a/tzrec/prompt/frontend_test.py +++ b/tzrec/prompt/frontend_test.py @@ -15,9 +15,8 @@ import numpy as np import torch -from tzrec.prompt import frontend -from tzrec.prompt.assembler import PromptAssembler as ReferenceAssembler -from tzrec.prompt.frontend import PromptAssembler, PromptFrontEnd, SlotTable, mix64 +from tzrec.prompt.assembler import HOLE_KEYS, PromptAssembler, host_lengths_key +from tzrec.prompt.frontend import SLOT_EMBEDS, PromptFrontEnd, SlotTable from tzrec.prompt.types import ( FillMode, PromptPlan, @@ -53,32 +52,28 @@ def _sid_space() -> ResolvedSidSpace: ) -def _slot(slot_id, name, fill, sequence=True, members=None, width=None): +def _slot(slot_id, name, fill): return SlotSeg( slot_id=slot_id, name=name, - feature_names=tuple(members or (name,)), - group_type=( - FeatureGroupType.JAGGED_SEQUENCE if sequence else FeatureGroupType.DEEP - ), - output_key=".sequence" if sequence else "", + feature_names=(name,), + group_type=FeatureGroupType.JAGGED_SEQUENCE, + output_key=".sequence", fill=fill, - width=width - or (Width(WidthKind.BOUNDED, 30) if sequence else Width(WidthKind.STATIC, 1)), + width=Width(WidthKind.BOUNDED, 30), ) -def _plan(segments, response_segments=(), projected=()): +def _plan(segments, projected=()): return PromptPlan( segments=segments, - response_segments=response_segments, + response_segments=(), max_length=256, max_total_length=None, max_holes=30, logits_suffix_len=4, - static_prefix_len=2, + static_prefix_len=1, projected_slots=projected, - slot_index={seg.name: i for i, seg in enumerate(projected)}, ) @@ -86,241 +81,6 @@ def _tensors(raw): return {key: torch.from_numpy(np.asarray(value)) for key, value in raw.items()} -class MixTest(unittest.TestCase): - def test_mix64_matches_a_host_reference(self): - """The mixer's masked shifts must reproduce SplitMix64 on negatives.""" - - def reference(value: int) -> int: - mask = (1 << 64) - 1 - z = (value * 0x9E3779B97F4A7C15) & mask - z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & mask - z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & mask - z = z ^ (z >> 31) - return z - (1 << 64) if z >= 1 << 63 else z - - values = [0, 1, -1, 2**31, -(2**31), 123456789, -987654321] - got = mix64(torch.tensor(values, dtype=torch.int64)).tolist() - self.assertEqual(got, [reference(v) for v in values]) - - def test_output_keys_match_the_module_constants(self): - """``forward`` writes literals; they must equal the exported names.""" - module = PromptAssembler(_plan((Static((7,)),))) - keys = set(module({"batch_size": torch.tensor(2)}).keys()) - self.assertEqual( - keys, - { - frontend.OUT_INPUT_IDS, - frontend.OUT_CU_SEQLENS, - frontend.OUT_HOLE_POSITIONS, - frontend.OUT_HOLE_KEYS, - frontend.OUT_HOLE_SLOT_COUNTS, - frontend.OUT_RESPONSE_LENGTHS, - }, - ) - - -class WalkTest(unittest.TestCase): - def setUp(self): - self.hist = _slot(0, "hist", FillMode.INLINE) - self.beh = _slot(1, "beh", FillMode.PROJECTED) - self.answer = _slot( - 2, "answer", FillMode.INLINE, width=Width(WidthKind.STATIC, 3) - ) - self.plan = _plan( - segments=(Static((10, 11)), self.hist, Static((12,)), self.beh), - response_segments=(self.answer,), - projected=(self.beh,), - ) - # a SID history as a multi-value sequence: items per row, codes per item - self.raw = { - "hist.values": np.array([0, 4, 8, 1, 5, 9, 2, 6, 10], dtype=np.int64), - "hist.lengths": np.array([2, 1], dtype=np.int64), - "hist.key_lengths": np.array([3, 3, 3], dtype=np.int64), - "beh.values": np.array([7, 8, 9, 21, 22], dtype=np.int64), - "beh.lengths": np.array([3, 2], dtype=np.int64), - "answer.values": np.array([3, 7, 11, 0, 4, 8], dtype=np.int64), - "answer.lengths": np.array([1, 1], dtype=np.int64), - "answer.key_lengths": np.array([3, 3], dtype=np.int64), - } - self.module = PromptAssembler( - self.plan, - _sid_space(), - features_are_multi_valued={"hist": True, "answer": True}, - plan_hash="a1b2c3d4e5f60718", - ) - - def _assert_matches_reference(self, module, raw): - out = module(_tensors(raw)) - reference = ReferenceAssembler(self.plan, _sid_space()).forward(raw) - for key, reference_key in ( - ("input_ids", "prompt_input_ids"), - ("cu_seqlens", "prompt_cu_seqlens"), - ("hole_positions", "prompt_hole_positions"), - ("response_lengths", "prompt_response_lengths"), - ): - self.assertEqual(out[key].tolist(), reference[reference_key].tolist(), key) - return out - - def test_walk_matches_the_reference_implementation(self): - """The tensor walk and the collator's walk are one specification.""" - self._assert_matches_reference(self.module, self.raw) - - def test_walk_matches_the_reference_on_the_flat_layout(self): - """A history stored one code per position walks the same way.""" - raw = { - "hist.values": self.raw["hist.values"], - "hist.lengths": np.array([6, 3], dtype=np.int64), - "beh.values": self.raw["beh.values"], - "beh.lengths": self.raw["beh.lengths"], - "answer.values": self.raw["answer.values"], - "answer.lengths": np.array([3, 3], dtype=np.int64), - } - module = PromptAssembler(self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") - self._assert_matches_reference(module, raw) - - def test_holes_land_on_sentinels(self): - """Every recorded hole is a sentinel and every sentinel is recorded.""" - out = self.module(_tensors(self.raw)) - sentinel = _sid_space().sentinel_token_id - holes = out["hole_positions"] - self.assertTrue(bool(torch.all(out["input_ids"][holes] == sentinel))) - self.assertEqual( - int(torch.sum(out["input_ids"] == sentinel)), int(holes.numel()) - ) - self.assertEqual(out["hole_slot_counts"].tolist(), [5]) - - def test_scripting_preserves_every_output(self): - """The artifact and the eager module are the same function.""" - batch = _tensors(self.raw) - eager = self.module(batch) - scripted = torch.jit.script(self.module)(batch) - for key, value in eager.items(): - self.assertTrue(torch.equal(scripted[key], value), key) - - def test_inline_without_a_sid_space_is_rejected(self): - with self.assertRaisesRegex(ValueError, "no sid_space"): - PromptAssembler(self.plan) - - -class FoldTest(unittest.TestCase): - def setUp(self): - self.beh = _slot(0, "beh", FillMode.PROJECTED) - self.plan = _plan(segments=(self.beh,), projected=(self.beh,)) - self.module = PromptAssembler( - self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718" - ) - - def _keys(self, values, lengths): - return self.module( - _tensors( - { - "beh.values": np.array(values, dtype=np.int64), - "beh.lengths": np.array(lengths, dtype=np.int64), - } - ) - )["hole_keys"] - - def test_equal_content_folds_equal(self): - """A key is a function of the hole's inputs and nothing else.""" - self.assertTrue( - torch.equal(self._keys([5, 6, 7], [3]), self._keys([5, 6, 7], [3])) - ) - - def test_different_content_folds_apart(self): - """The whole point: a changed input must not reuse the cached KV.""" - first = self._keys([5, 6, 7], [3]) - second = self._keys([5, 6, 8], [3]) - self.assertEqual(first[:2].tolist(), second[:2].tolist()) - self.assertNotEqual(int(first[2]), int(second[2])) - - def test_the_plan_hash_salts_the_keys(self): - """Two plans must not cross-match in a shared prefix cache.""" - other = PromptAssembler(self.plan, _sid_space(), plan_hash="ffffffffffffffff") - batch = _tensors( - { - "beh.values": np.array([5, 6, 7], dtype=np.int64), - "beh.lengths": np.array([3], dtype=np.int64), - } - ) - self.assertFalse( - torch.equal(self.module(batch)["hole_keys"], other(batch)["hole_keys"]) - ) - - def test_a_permuted_multi_value_item_does_not_collide(self): - """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" - module = PromptAssembler( - self.plan, - _sid_space(), - features_are_multi_valued={"beh": True}, - plan_hash="a1b2c3d4e5f60718", - ) - - def keys(values): - return module( - _tensors( - { - "beh.values": np.array(values, dtype=np.int64), - "beh.lengths": np.array([1], dtype=np.int64), - "beh.key_lengths": np.array([3], dtype=np.int64), - } - ) - )["hole_keys"] - - self.assertNotEqual(keys([1, 5, 9]).tolist(), keys([9, 5, 1]).tolist()) - - def test_two_members_exchanging_values_do_not_collide(self): - """Without the member index a two-member slot is order-blind.""" - slot = _slot(0, "pair", FillMode.PROJECTED, members=("a", "b")) - plan = _plan(segments=(slot,), projected=(slot,)) - module = PromptAssembler(plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") - - def keys(first, second): - return module( - _tensors( - { - "a.values": np.array(first, dtype=np.int64), - "a.lengths": np.array([1], dtype=np.int64), - "b.values": np.array(second, dtype=np.int64), - "b.lengths": np.array([1], dtype=np.int64), - } - ) - )["hole_keys"] - - self.assertNotEqual(keys([3], [9]).tolist(), keys([9], [3]).tolist()) - - def test_a_dense_member_folds_its_bit_pattern(self): - """A float member contributes the parsed input verbatim, so it is stable.""" - slot = _slot(0, "vec", FillMode.PROJECTED, sequence=False) - plan = _plan(segments=(slot,), projected=(slot,)) - module = PromptAssembler( - plan, _sid_space(), features_are_dense={"vec": True}, plan_hash="a1b2" - ) - - def keys(rows): - return module( - { - "vec.values": torch.tensor(rows, dtype=torch.float32), - "vec.lengths": torch.ones(len(rows), dtype=torch.int64), - } - )["hole_keys"] - - self.assertEqual(keys([[0.5, 1.0]]).tolist(), keys([[0.5, 1.0]]).tolist()) - self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) - - @unittest.skipIf(not torch.cuda.is_available(), "no GPU") - def test_the_fold_is_bit_identical_across_devices(self): - """Integer addition cannot depend on the order a device reduces in.""" - batch = _tensors( - { - "beh.values": np.arange(64, dtype=np.int64), - "beh.lengths": np.array([32, 32], dtype=np.int64), - } - ) - on_cpu = self.module(batch)["hole_keys"] - on_gpu = self.module({k: v.cuda() for k, v in batch.items()})["hole_keys"] - self.assertTrue(torch.equal(on_cpu, on_gpu.cpu())) - - class FrontEndTest(unittest.TestCase): def setUp(self): self.slot = _slot(0, "beh", FillMode.PROJECTED) @@ -353,7 +113,7 @@ def _front_end(self, tables, embed_keys=(("beh",),)): def test_in_module_tables_produce_one_embedding_per_hole(self): """Without a host lookup stage, the artifact carries the tables.""" out = self._front_end([self._table()])(self.batch) - self.assertEqual(tuple(out["slot_embeds"].shape), (3, 16)) + self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (3, 16)) self.assertEqual(int(out["hole_positions"].numel()), 3) def test_both_lookup_shapes_agree_given_the_same_rows(self): @@ -365,15 +125,15 @@ def test_both_lookup_shapes_agree_given_the_same_rows(self): batch = { "beh.values": self.batch["beh.values"], "beh": table(self.batch), - frontend.host_lengths_key("beh"): self.batch["beh.lengths"], + host_lengths_key("beh"): self.batch["beh.lengths"], } self.assertTrue( torch.allclose( - with_tables(self.batch)["slot_embeds"], host(batch)["slot_embeds"] + with_tables(self.batch)[SLOT_EMBEDS], host(batch)[SLOT_EMBEDS] ) ) self.assertTrue( - torch.equal(with_tables(self.batch)["hole_keys"], host(batch)["hole_keys"]) + torch.equal(with_tables(self.batch)[HOLE_KEYS], host(batch)[HOLE_KEYS]) ) def test_the_device_argument_is_optional(self): @@ -394,7 +154,7 @@ def test_the_artifact_reloads_with_its_identity(self): ("vocab", "plan", "test-bundle"), ) out = loaded(self.batch) - self.assertEqual(tuple(out["slot_embeds"].shape), (3, 16)) + self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (3, 16)) def test_pattern_i_exports_empty_hole_streams(self): """With no projected slot the artifact degenerates, and still exists.""" @@ -413,7 +173,7 @@ def test_pattern_i_exports_empty_hole_streams(self): ) self.assertEqual(out["input_ids"].tolist(), [10, 1000, 1005, 1010]) self.assertEqual(int(out["hole_positions"].numel()), 0) - self.assertEqual(tuple(out["slot_embeds"].shape), (0, 0)) + self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (0, 0)) if __name__ == "__main__": diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index efdd7256..e61b1486 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -185,7 +185,6 @@ class PromptPlan: projected_slots: PROJECTED occurrences in emission order, which is also ascending hole position; nothing may reorder them by slot id or by shared module, because the serving scatter is positional. - slot_index: slot name to its index in ``projected_slots``. fold: the constants ``hole_keys`` mixes in. """ @@ -197,7 +196,6 @@ class PromptPlan: logits_suffix_len: Optional[int] static_prefix_len: int projected_slots: Tuple[SlotSeg, ...] - slot_index: Mapping[str, int] = field(default_factory=dict) fold: FoldConstants = field(default_factory=FoldConstants) diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index ff97211d..5d74c52d 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -27,7 +27,7 @@ import torch from safetensors.torch import load_file -from tzrec.prompt.frontend import mix64 +from tzrec.prompt.assembler import mix64 from tzrec.prompt.types import FoldConstants from tzrec.tests.prompt_test_util import export_tiny_genrec, offset_sid_codes from tzrec.utils.test_util import make_test_dir diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index 42278a42..7aa3046f 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -21,7 +21,13 @@ from tzrec.datasets.utils import Batch from tzrec.main import export from tzrec.models.model import TrainWrapper -from tzrec.prompt.assembler import PromptAssembler +from tzrec.prompt.assembler import ( + CU_SEQLENS, + HOLE_KEYS, + HOLE_POSITIONS, + INPUT_IDS, + PromptAssembler, +) from tzrec.tests.prompt_test_util import ( GenrecModelTestBase, assemble_into, @@ -44,11 +50,8 @@ def _batch_from_codes(self, hist, answer): "answer.values": torch.tensor(offset_sid_codes(answer, _CODEBOOK)), "answer.lengths": torch.tensor([len(answer)]), } - streams = assemble_into(self.compiled_prompt, parsed) batch = Batch() - batch.additional_infos.update( - {k: torch.from_numpy(np.asarray(v)) for k, v in streams.items()} - ) + batch.additional_infos.update(assemble_into(self.compiled_prompt, parsed)) return batch def test_written_digests_satisfy_the_restore_guard(self) -> None: @@ -122,25 +125,21 @@ def test_export_writes_a_loadable_serving_directory(self) -> None: os.path.exists(os.path.join(exported.export_dir, "frontend/sparse")) ) - # the collator and the scripted front-end are two call sites of one walk + # the collator and the exported artifact are two call sites of one walk batch = _serving_batch() compiled = exported.compiled_prompt collator = PromptAssembler( - compiled.prompt_plan, compiled.sid_space, include_response=False - ).forward({k: v.numpy() for k, v in batch.items()}) + compiled.prompt_plan, + compiled.sid_space, + plan_hash=compiled.plan_hash, + include_response=False, + )(batch) front_end = torch.jit.load( os.path.join(exported.export_dir, "frontend", "scripted_model.pt") ) out = front_end(batch, torch.device("cpu")) - self.assertEqual( - out["input_ids"].tolist(), collator["prompt_input_ids"].tolist() - ) - self.assertEqual( - out["cu_seqlens"].tolist(), collator["prompt_cu_seqlens"].tolist() - ) - self.assertEqual( - out["hole_positions"].tolist(), collator["prompt_hole_positions"].tolist() - ) + for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, HOLE_KEYS): + self.assertTrue(torch.equal(out[key], collator[key]), key) self.assertEqual(tuple(out["slot_embeds"].shape), (2, 32)) def test_distributed_embedding_export_writes_the_processor_shape(self) -> None: diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index 396aca96..b15697d8 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -24,7 +24,7 @@ from tzrec.features.feature import BaseFeature, FgMode, create_features from tzrec.main import _create_features, _create_model, export from tzrec.models.model import TrainWrapper -from tzrec.prompt.assembler import PromptAssembler +from tzrec.prompt.assembler import PROMPT_INFO_PREFIX, PromptAssembler from tzrec.prompt.compile import compile_prompt from tzrec.prompt.types import CompiledPrompt from tzrec.protos import feature_pb2 @@ -86,22 +86,25 @@ def offset_sid_codes(codes: Sequence[Any], codebook: Sequence[int]) -> np.ndarra def assemble_into( - compiled_prompt: CompiledPrompt, - parsed_features: Dict[str, "np.ndarray"], -) -> Dict[str, np.ndarray]: - """Assemble one parsed batch with a temporary assembler. + compiled_prompt: CompiledPrompt, parsed_features: Dict[str, Any] +) -> Dict[str, torch.Tensor]: + """Assemble one parsed batch the way the collator does. Args: compiled_prompt: the compiled prompt. parsed_features: ``{column}.values`` / ``{column}.lengths`` as the data - parser emits them. + parser emits them, as tensors or arrays. Returns: The assembled streams keyed for ``additional_infos``. """ - return PromptAssembler( - compiled_prompt.prompt_plan, compiled_prompt.sid_space - ).forward(parsed_features) + batch = {k: torch.as_tensor(np.asarray(v)) for k, v in parsed_features.items()} + streams = PromptAssembler( + compiled_prompt.prompt_plan, + compiled_prompt.sid_space, + plan_hash=compiled_prompt.plan_hash, + )(batch) + return {PROMPT_INFO_PREFIX + k: v for k, v in streams.items()} _CODEBOOK = [4, 4, 4] @@ -165,9 +168,7 @@ def _model( def _batch(self, parsed, compiled_prompt=None, sparse=None): streams = assemble_into(compiled_prompt or self.compiled_prompt, parsed) batch = Batch(sparse_features={BASE_DATA_GROUP: sparse} if sparse else {}) - batch.additional_infos.update( - {k: torch.from_numpy(np.asarray(v)) for k, v in streams.items()} - ) + batch.additional_infos.update(streams) return batch From 40169c3be6a97e9e167628a1c1069ddf757147e2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 17:31:04 +0800 Subject: [PATCH 03/21] [refactor] export the genrec front-end through export_model The serving front-end was assembled by hand: slot tables copied out of the restored EmbeddingGroup, a bespoke module, hand-written dense_meta.json and sparse npz writers, a frontend/ subdirectory and a class-specific block in export(), all re-implementing what export_model already does for every model. The served half of a genrec model is now GenrecFrontEnd, a TowerWrapper-style wrapper over the model's own embedding group and projections, handed to export_model beside MatchModel and TDM. ScriptWrapper assembles the prompt from the parsed dict as the collator does, the assembler is an FX leaf so tracing keeps it whole for TorchScript, and the HuggingFace weights, composite config and prompt/ are written by an export_assets hook the export discovers the way checkpointing discovers hf_backbone. Quantization, INPUT_TILE and the distributed-embedding split therefore apply unchanged, and the export is one flat directory. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- docs/source/usage/export.md | 15 +- tzrec/datasets/dataset.py | 8 +- tzrec/main.py | 111 +---- tzrec/main_test.py | 16 +- tzrec/models/genrec_model.py | 234 ++++++++++- tzrec/models/genrec_model_test.py | 109 ++++- tzrec/models/model.py | 24 +- tzrec/prompt/assembler.py | 13 + tzrec/prompt/export.py | 422 -------------------- tzrec/prompt/export_test.py | 249 ------------ tzrec/prompt/frontend.py | 215 ---------- tzrec/prompt/frontend_test.py | 180 --------- tzrec/tests/genrec_serving_contract_test.py | 25 +- tzrec/tests/prompt_integration_test.py | 133 +++--- tzrec/tests/prompt_test_util.py | 99 ++++- tzrec/utils/export_util.py | 34 +- tzrec/utils/fx_util.py | 8 +- tzrec/utils/hf_export_util.py | 38 +- 18 files changed, 649 insertions(+), 1284 deletions(-) delete mode 100644 tzrec/prompt/export.py delete mode 100644 tzrec/prompt/export_test.py delete mode 100644 tzrec/prompt/frontend.py delete mode 100644 tzrec/prompt/frontend_test.py diff --git a/docs/source/usage/export.md b/docs/source/usage/export.md index da38097f..0369ba1c 100644 --- a/docs/source/usage/export.md +++ b/docs/source/usage/export.md @@ -226,22 +226,23 @@ eascmd -i ${ACCESS_KEY_ID} -k ${ACCESS_KEY_SECRET} -e ${ENDPOINT} create aot_exp ## 生成式推荐模型(genrec)导出 -`genrec_causal_lm_model` 导出为一个 HuggingFace 目录,供 SGLang 等 LLM 推理引擎直接加载,并附带在线 Processor 所需的 prompt 前端: +`genrec_causal_lm_model` 与其他模型一样通过 `tzrec.export` 导出:tzrec 侧导出的是 **prompt 前端**(特征 -> 拼接后的 token 流与 PROJECTED slot 的投影向量),量化、`INPUT_TILE`、`USE_DISTRIBUTED_EMBEDDING` 等环境变量与普通模型一致;LLM 骨干网络则以 HuggingFace 权重的形式放在同一目录下,交给 SGLang 等 LLM 推理引擎解码: ``` export_dir/ + scripted_model.pt # prompt 前端:特征 -> input_ids / hole_positions / slot_embeds / hole_keys / hole_slot_counts + fg.json pipeline.config model_acc.json + dense_meta.json sparse/ # 仅 USE_DISTRIBUTED_EMBEDDING=1 时,与普通模型的分布式 embedding 导出一致 config.json # 复合结构:architectures 为 PromptGenRecForCausalLM,骨干网络配置位于 text_config model.safetensors # 骨干网络权重,参数名与骨干网络一致 generation_config.json prompt/ - prompt.json # 服务契约:sid_space(band、base_vocab_size、bundle_uuid)、decode 调度、vocab_hash/plan_hash、frontend 说明 + prompt.json # 服务契约:sid_space(band、base_vocab_size、bundle_uuid)、decode 调度、vocab_hash/plan_hash、frontend 输入输出 tokenizer/ # 扩展了 SID token 的 tokenizer,可由 AutoTokenizer 加载(对应 SGLang 的 --tokenizer-path) - frontend/ # 标准 tzrec 模型目录,TorchEasyRec Processor 以 model_path 直接加载 - scripted_model.pt # prompt 前端:特征 -> input_ids / hole_positions / slot_embeds / hole_keys / hole_slot_counts - fg.json pipeline.config model_acc.json ``` -- 默认导出下,PROJECTED slot 的 embedding 表内置于 `frontend/scripted_model.pt`,前端输入为 FG 输出的原始 id(`{feature}.values` / `.lengths` / `.key_lengths`)。 -- 设置 `USE_DISTRIBUTED_EMBEDDING=1` 时,前端改为读取 Processor 分布式 embedding 阶段查表后的向量(命名由 `frontend/dense_meta.json` 描述),embedding 表以 `frontend/sparse/*.npz` 导出,格式与普通模型的分布式 embedding 导出一致。此时前端仍需要 INLINE slot 与 PROJECTED 成员的原始 id(用于计算 `hole_keys`),`prompt.json` 的 `frontend.inputs` 列出了全部输入 key。 +- TorchEasyRec Processor 以 `export_dir` 作为 `model_path` 加载前端;SGLang 以 `--model-path export_dir` 加载骨干网络,并通过 `SGLANG_PROMPT_FRONTEND_PATH=export_dir/scripted_model.pt` 加载前端。 +- 默认导出下前端内置 PROJECTED slot 的 embedding 表,输入为 FG 输出的原始 id(`{feature}.values` / `.lengths` / `.key_lengths`)。设置 `USE_DISTRIBUTED_EMBEDDING=1` 时,前端改为读取 Processor 分布式 embedding 阶段查表后的向量(命名由 `dense_meta.json` 描述),embedding 表以 `sparse/*.npz` 导出。此时前端仍需要 INLINE slot 与 PROJECTED 成员的原始 id(用于计算 `hole_keys`),`prompt.json` 的 `frontend.inputs` 列出了全部输入 key。 +- 前端仅支持 TorchScript 导出:prompt 拼接的形状随请求变化,`ENABLE_AOT` / `ENABLE_TRT` / `USE_RTP` 不适用。 - 约束解码索引不由 tzrec 生成:推理侧根据 `prompt/prompt.json` 与 SID bundle 的 `sid_to_items` 构建(SGLang 侧 `python -m sglang.srt.beam_search.build_constraint_csr`),并以 `bundle_uuid` 校验索引与模型是否来自同一 bundle。 - genrec 导出不支持 `export_config.use_dense_ema=true`。 diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index 33b48235..bc851322 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -40,7 +40,7 @@ remove_nullable, ) from tzrec.features.feature import BaseFeature -from tzrec.prompt.assembler import PROMPT_INFO_PREFIX, PromptAssembler +from tzrec.prompt.assembler import OUTPUT_KEYS, PROMPT_INFO_PREFIX, PromptAssembler from tzrec.prompt.types import CompiledPrompt from tzrec.protos import data_pb2 from tzrec.utils import config_util @@ -401,11 +401,9 @@ def _build_batch(self, input_data: Dict[str, pa.Array]) -> Batch: batch = self._data_parser.to_batch(output_data) if self._prompt_assembler is not None: + streams = self._prompt_assembler(output_data) batch.additional_infos.update( - { - PROMPT_INFO_PREFIX + k: v - for k, v in self._prompt_assembler(output_data).items() - } + {PROMPT_INFO_PREFIX + k: streams[k] for k in OUTPUT_KEYS} ) # Set checkpoint info on batch diff --git a/tzrec/main.py b/tzrec/main.py index a841d8b1..c469e173 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -52,7 +52,7 @@ BaseFeature, create_features, ) -from tzrec.models.genrec_model import BaseGenrecModel +from tzrec.models.genrec_model import BaseGenrecModel, GenrecFrontEnd from tzrec.models.match_model import ( MatchModel, MatchTower, @@ -74,12 +74,7 @@ from tzrec.optim.lr_scheduler import BaseLR from tzrec.optim.optimizer import TZRecOptimizer from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.persist import ( - PROMPT_DIR, - TOKENIZER_DIR, - check_prompt_assets, - write_serving_contract, -) +from tzrec.prompt.persist import check_prompt_assets from tzrec.prompt.types import CompiledPrompt from tzrec.protos.data_pb2 import DataConfig, DatasetType from tzrec.protos.eval_pb2 import EvalConfig @@ -1160,52 +1155,6 @@ def evaluate( logger.info("Evaluate Finished.") -def _export_prompt_serving_assets( - pipeline_config: EasyRecConfig, - features: List[BaseFeature], - compiled_prompt: CompiledPrompt, - checkpoint_path: str, - export_dir: str, -) -> None: - """Write the composite config, the front-end and the serving contract. - - The model is rebuilt and restored on CPU so the front-end carries the - trained projections and slot tables by reference, the same way the - TorchScript export restores a checkpoint before scripting it. - """ - from tzrec.prompt.export import ( - build_front_end, - export_sparse_tables, - write_composite_config, - write_front_end_dir, - ) - from tzrec.utils.state_dict_util import init_parameters - - write_composite_config(export_dir) - - model = _create_model( - pipeline_config.model_config, - features, - list(pipeline_config.data_config.label_fields), - compiled_prompt=compiled_prompt, - ) - model.set_is_inference(True) - wrapped_model = ScriptWrapper(model) - init_parameters(wrapped_model, torch.device("cpu")) - checkpoint_util.restore_model(checkpoint_path, wrapped_model) - - carry_tables = not acc_utils.use_distributed_embedding() - front_end, frontend_meta, dense_meta = build_front_end( - model, compiled_prompt, features, carry_tables - ) - frontend_dir = write_front_end_dir( - front_end, pipeline_config, features, export_dir, dense_meta - ) - if not carry_tables: - export_sparse_tables(wrapped_model, checkpoint_path, frontend_dir) - write_serving_contract(compiled_prompt, export_dir, frontend_meta) - - def export( pipeline_config_path: str, export_dir: str, @@ -1258,57 +1207,22 @@ def export( else: checkpoint_path, _ = ckpt_manager.latest_checkpoint() - model_cls = _get_model_class(pipeline_config.model_config) - if issubclass(model_cls, BaseGenrecModel): - if config_util.use_dense_ema( - pipeline_config.export_config, pipeline_config.train_config - ): - raise ValueError( - "HF export: dcp_to_hf reads /model, so it cannot " - "serve Dense EMA parameters. Set export_config.use_dense_ema to " - "false to export the raw weights." - ) - if not checkpoint_path: - raise ValueError("HF export: no checkpoint found to convert.") - if not os.path.exists(os.path.join(checkpoint_path, "config.json")): - raise ValueError( - f"HF export: {checkpoint_path} has no co-located HF assets; it " - f"was not written by an HF-backed model." - ) - if assets: - logger.warning(f"HF export ignores asset_files: {assets}.") - features = _create_features( - list(pipeline_config.feature_configs), pipeline_config.data_config - ) - compiled_prompt = compile_prompt( - pipeline_config.prompt_config, - features, - list(pipeline_config.data_config.label_fields), - tokenizer_dir=os.path.join(export_dir, PROMPT_DIR, TOKENIZER_DIR) - if is_rank_zero - else None, - ) - check_prompt_assets(compiled_prompt, checkpoint_path) - if is_rank_zero: - from tzrec.utils.hf_export_util import dcp_to_hf - - dcp_to_hf(checkpoint_path, export_dir) - _export_prompt_serving_assets( - pipeline_config, features, compiled_prompt, checkpoint_path, export_dir - ) - return - data_config = pipeline_config.data_config # Build feature features = _create_features(list(pipeline_config.feature_configs), data_config) + compiled_prompt = _compile_prompt(pipeline_config, features) + if checkpoint_path: + check_prompt_assets(compiled_prompt, checkpoint_path) + # Build model model = _create_model( pipeline_config.model_config, features, list(data_config.label_fields), sampler_type=None, + compiled_prompt=compiled_prompt, ) InferWrapper = ScriptWrapper # Flip to inference *before* wrapping so view-dependent state @@ -1356,6 +1270,17 @@ def export( os.path.join(export_dir, "model"), assets=assets, ) + elif isinstance(model.model, BaseGenrecModel): + # tzrec serves the prompt front-end; the LM rides beside it as + # HuggingFace weights for the engine that decodes + export_model( + ori_pipeline_config, + InferWrapper(GenrecFrontEnd(model.model)), + checkpoint_path, + export_dir, + assets=assets, + additional_export_config=additional_export_config, + ) else: export_model( ori_pipeline_config, diff --git a/tzrec/main_test.py b/tzrec/main_test.py index 3b7196e8..1359ec60 100644 --- a/tzrec/main_test.py +++ b/tzrec/main_test.py @@ -20,14 +20,12 @@ import pyarrow as pa import torch -from google.protobuf import text_format from parameterized import parameterized from tzrec.datasets.utils import RecordBatchTensor from tzrec.main import ( _create_model, _train_and_evaluate, - export, predict, predict_checkpoint, ) @@ -369,19 +367,11 @@ def test_train_loop_pairs_dense_export_with_delta_dump(self) -> None: def test_hf_export_rejects_dense_ema(self) -> None: # dcp_to_hf reads /model unconditionally, so it would silently # ship raw weights where TorchScript export ships the EMA ones. + from tzrec.tests.prompt_test_util import export_tiny_genrec + with tempfile.TemporaryDirectory() as test_dir: - config = EasyRecConfig() - config.train_input_path = "unused" - config.eval_input_path = "unused" - config.model_dir = os.path.join(test_dir, "train") - os.makedirs(config.model_dir) - config.train_config.dense_optimizer.ema.CopyFrom(EMAConfig()) - config.model_config.genrec_causal_lm_model.SetInParent() - config_path = os.path.join(test_dir, "pipeline.config") - with open(config_path, "w") as f: - f.write(text_format.MessageToString(config)) with self.assertRaisesRegex(ValueError, "Dense EMA"): - export(config_path, os.path.join(test_dir, "export")) + export_tiny_genrec(test_dir, use_dense_ema=True) class PredictionLifecycleTest(unittest.TestCase): diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 7e309cda..87a9092c 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -15,30 +15,55 @@ projections, converts SID coordinate systems, scores the response window and supplies the digests a checkpoint records. A family subclass owns its forward and decode path. + +``GenrecFrontEnd`` is the half of the model tzrec serves: the assembled prompt +and the projected slots, everything before the LM's embedding gather. It is +exported like any tzrec model; the LM itself is handed to an LLM engine as the +HuggingFace weights beside it. """ import inspect -from typing import Any, Dict, List, Optional, Tuple +import os +from typing import Any, Dict, List, Optional, Sequence, Tuple import torch import torchmetrics from torch import nn from transformers import AutoConfig, AutoModelForCausalLM +from tzrec.acc import utils as acc_utils from tzrec.datasets.utils import Batch from tzrec.features.feature import BaseFeature from tzrec.models.model import BaseModel from tzrec.modules.embedding import EmbeddingGroup from tzrec.modules.prompt_projection import PromptProjection from tzrec.prompt.assembler import ( + CU_SEQLENS, + HOLE_KEYS, + HOLE_POSITIONS, + HOLE_SLOT_COUNTS, + INPUT_IDS, + PROMPT_CU_SEQLENS, + PROMPT_HOLE_KEYS, PROMPT_HOLE_POSITIONS, + PROMPT_HOLE_SLOT_COUNTS, PROMPT_INPUT_IDS, ) -from tzrec.prompt.types import CompiledPrompt -from tzrec.protos.model_pb2 import ModelConfig +from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.persist import PROMPT_DIR, TOKENIZER_DIR, write_serving_contract +from tzrec.prompt.types import CompiledPrompt, FillMode, PromptPlan, SlotSeg +from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.protos.pipeline_pb2 import EasyRecConfig +from tzrec.utils import config_util, env_util +from tzrec.utils.hf_export_util import dcp_to_hf, write_composite_config from tzrec.utils.logging_util import logger +SLOT_EMBEDS = "slot_embeds" +SCRIPTED_MODEL_FILENAME = "scripted_model.pt" +LOOKUP_ARTIFACT = "artifact" +LOOKUP_HOST = "host" + _PARAM_DTYPE: Dict[int, torch.dtype] = { GenrecModelConfig.FP32: torch.float32, GenrecModelConfig.BF16: torch.bfloat16, @@ -188,6 +213,11 @@ def hf_backbone(self) -> nn.Module: """The HF module export and checkpointing reach for.""" return self.lm + @property + def compiled_prompt(self) -> CompiledPrompt: + """The prompt this model was built against.""" + return self._prompt + def build_input(self, batch: Batch) -> torch.Tensor: """Build packed LM input embeddings and fill projected positions. @@ -199,27 +229,18 @@ def build_input(self, batch: Batch) -> torch.Tensor: """ ids = batch.additional_infos[PROMPT_INPUT_IDS] embeds = self.lm.get_input_embeddings()(ids) - - prompt_plan = self._prompt.prompt_plan - if not prompt_plan.projected_slots: + if not self._prompt.prompt_plan.projected_slots: return embeds - - grouped = self.embedding_group(batch) - hidden_size = embeds.shape[-1] - projected_embeddings = [ - proj(grouped[seg.name + seg.output_key]).reshape(-1, hidden_size) - for seg, proj in zip(prompt_plan.projected_slots, self._slot_projections) - ] - # The assembler records holes in this projected-occurrence-major order. + projected = project_slots( + self.embedding_group, + self._prompt.prompt_plan, + self._slot_projections, + batch, + embeds.shape[-1], + ) # out of place: embeds carries grad from the embedding lookup return embeds.index_copy( - 0, - batch.additional_infos[PROMPT_HOLE_POSITIONS], - ( - projected_embeddings[0] - if len(projected_embeddings) == 1 - else torch.cat(projected_embeddings) - ).to(embeds.dtype), + 0, batch.additional_infos[PROMPT_HOLE_POSITIONS], projected.to(embeds.dtype) ) def _tokens_to_local_codes( @@ -311,3 +332,174 @@ def init_from_pretrained(self) -> None: ) self.lm.load_state_dict(pretrained.state_dict()) del pretrained + + +def project_slots( + embedding_group: EmbeddingGroup, + prompt_plan: PromptPlan, + slot_projections: Sequence[nn.Module], + batch: Batch, + hidden_size: int, +) -> torch.Tensor: + """Look every projected slot up and project it into the LM input space. + + Args: + embedding_group: the model's prompt groups. + prompt_plan: fixes the slot order. + slot_projections: one module per projected slot, in the same order. + batch: the batch to look up. + hidden_size: the LM hidden size. + + Returns: + ``(total_holes, hidden_size)`` in the order the assembler records holes: + projected occurrence first, then sample. + """ + grouped = embedding_group(batch) + parts = [ + proj(grouped[seg.name + seg.output_key]).reshape(-1, hidden_size) + for seg, proj in zip(prompt_plan.projected_slots, slot_projections) + ] + return parts[0] if len(parts) == 1 else torch.cat(parts) + + +class GenrecFrontEnd(nn.Module): + """The served half of a genrec model, exported like any tzrec model. + + It shares the model's embedding group and projections, so under the + inference wrapper their state-dict names are the checkpoint's and the LM is + never loaded at export. ``predict`` returns the assembled prompt and the + projected slot embeddings; an LLM engine gathers the LM's own table, scatters + ``slot_embeds`` at ``hole_positions`` and decodes. + + Args: + model: the genrec model to serve. + """ + + def __init__(self, model: BaseGenrecModel) -> None: + super().__init__() + if acc_utils.is_aot() or acc_utils.is_trt() or env_util.use_rtp(): + raise ValueError( + "the genrec front-end is exported with TorchScript only: its " + "prompt walk has data-dependent shapes, which AOT, TRT and RTP " + "export cannot capture. Unset ENABLE_AOT / ENABLE_TRT / USE_RTP." + ) + self.embedding_group = model.embedding_group + self.projections = model.projections + self._slot_projections = list(model._slot_projections) + self._prompt = model.compiled_prompt + self._features = list(model.features) + self._hidden_size = int(model.lm.config.hidden_size) + + @property + def features(self) -> List[BaseFeature]: + """The features the served prompt reads.""" + return self._features + + @property + def feature_groups(self) -> List[FeatureGroupConfig]: + """The groups derived for the projected slots.""" + return list(self._prompt.projection_plan.feature_groups) + + @property + def compiled_prompt(self) -> CompiledPrompt: + """The prompt the inference wrapper assembles before ``predict``.""" + return self._prompt + + def predict(self, batch: Batch) -> Dict[str, torch.Tensor]: + """Return the assembled streams and the projected slot embeddings. + + Args: + batch: carries the assembled prompt in ``additional_infos``. + + Returns: + The serving contract's outputs. + """ + infos = batch.additional_infos + out = { + INPUT_IDS: infos[PROMPT_INPUT_IDS], + CU_SEQLENS: infos[PROMPT_CU_SEQLENS], + HOLE_POSITIONS: infos[PROMPT_HOLE_POSITIONS], + HOLE_KEYS: infos[PROMPT_HOLE_KEYS], + HOLE_SLOT_COUNTS: infos[PROMPT_HOLE_SLOT_COUNTS], + } + if self._prompt.prompt_plan.projected_slots: + out[SLOT_EMBEDS] = project_slots( + self.embedding_group, + self._prompt.prompt_plan, + self._slot_projections, + batch, + self._hidden_size, + ) + else: + out[SLOT_EMBEDS] = torch.zeros( + 0, 0, dtype=torch.float32, device=infos[PROMPT_INPUT_IDS].device + ) + return out + + def _serving_contract(self) -> Dict[str, Any]: + """What the exported module reads and writes, for ``prompt.json``.""" + by_name = {feature.name: feature for feature in self._features} + inputs: List[str] = [] + plan = self._prompt.prompt_plan + for seg in plan.segments: + if not isinstance(seg, SlotSeg): + continue + members = ( + seg.feature_names[:1] + if seg.fill is FillMode.INLINE + else seg.feature_names + ) + for name in members: + feature = by_name[name] + inputs.extend([f"{name}.values", f"{name}.lengths"]) + if feature.is_sequence and feature.value_dim != 1: + inputs.append(f"{name}.key_lengths") + return { + "model": SCRIPTED_MODEL_FILENAME, + "lookup": ( + LOOKUP_HOST + if acc_utils.use_distributed_embedding() + else LOOKUP_ARTIFACT + ), + "inputs": list(dict.fromkeys(inputs)), + "outputs": [ + INPUT_IDS, + CU_SEQLENS, + HOLE_POSITIONS, + HOLE_KEYS, + HOLE_SLOT_COUNTS, + SLOT_EMBEDS, + ], + } + + def export_assets( + self, pipeline_config: EasyRecConfig, checkpoint_path: str, save_dir: str + ) -> None: + """Write what an LLM engine reads beside the scripted front-end. + + The HuggingFace weights and composite config, the extended tokenizer and + the serving contract. Called by the export on rank 0, inside its save + dir. + + Args: + pipeline_config: the pipeline being exported. + checkpoint_path: the checkpoint the weights come from. + save_dir: the export directory. + """ + if config_util.use_dense_ema( + pipeline_config.export_config, pipeline_config.train_config + ): + raise ValueError( + "HF export: dcp_to_hf reads /model, so it cannot " + "serve Dense EMA parameters. Set export_config.use_dense_ema to " + "false to export the raw weights." + ) + dcp_to_hf(checkpoint_path, save_dir) + write_composite_config(save_dir) + compile_prompt( + pipeline_config.prompt_config, + self._features, + list(pipeline_config.data_config.label_fields), + tokenizer_dir=os.path.join(save_dir, PROMPT_DIR, TOKENIZER_DIR), + ) + write_serving_contract(self._prompt, save_dir, self._serving_contract()) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 39278856..0b1b170c 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -10,6 +10,9 @@ # limitations under the License. import dataclasses +import json +import os +import shutil import unittest import torch @@ -18,13 +21,24 @@ from transformers import AutoModelForCausalLM from tzrec.datasets.utils import Batch -from tzrec.models.genrec_model import _PARAM_DTYPE +from tzrec.models.genrec_model import ( + _PARAM_DTYPE, + SLOT_EMBEDS, + GenrecFrontEnd, + project_slots, +) +from tzrec.models.model import ScriptWrapper from tzrec.prompt.assembler import ( + HOLE_KEYS, + HOLE_POSITIONS, + INPUT_IDS, PROMPT_HOLE_POSITIONS, PROMPT_INPUT_IDS, + PromptAssembler, ) from tzrec.prompt.compile import compile_prompt from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig from tzrec.tests.prompt_test_util import ( _CODEBOOK, @@ -33,7 +47,9 @@ create_prompt_feature, offset_sid_codes, projected_feature, + write_genrec_checkpoint, ) +from tzrec.utils.fx_util import symbolic_trace from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( parameterized_name_func, @@ -211,5 +227,96 @@ def test_init_from_pretrained_replaces_the_empty_weights(self) -> None: torch.testing.assert_close(after, expected) +class GenrecFrontEndTest(GenrecModelTestBase): + """The served half of the model, under the same wrapper every export uses.""" + + def setUp(self) -> None: + super().setUp() + self.features = [ + create_prompt_feature(_HIST), + create_prompt_feature(projected_feature("beh", 8)), + ] + self.prompt_config = PromptConfig( + tokenizer_path=self.tok, + prompt="History : {{hist}} . {{beh}} Predict :", + response="{{answer}}", + ) + self.prompt_config.sid_space.codebook.extend(_CODEBOOK) + self.compiled_prompt = compile_prompt( + self.prompt_config, self.features, ["answer"] + ) + self.model = self._model() + init_parameters(self.model, device=torch.device("cpu")) + # the parsed dict as the data parser emits it: a dense sequence of codes + # and a sparse behaviour sequence + self.data = { + "hist.values": torch.tensor( + offset_sid_codes([0, 1, 2, 3, 0, 1], _CODEBOOK), dtype=torch.float32 + ).reshape(-1, 1), + "hist.lengths": torch.tensor([6]), + "beh.values": torch.tensor([3, 9]), + "beh.lengths": torch.tensor([2]), + } + self.wrapped = ScriptWrapper(GenrecFrontEnd(self.model)) + + def test_front_end_returns_the_walk_and_the_projected_slots(self) -> None: + out = self.wrapped(self.data) + compiled = self.compiled_prompt + walk = PromptAssembler( + compiled.prompt_plan, + compiled.sid_space, + plan_hash=compiled.plan_hash, + include_response=False, + )(self.data) + for key in (INPUT_IDS, HOLE_POSITIONS, HOLE_KEYS): + self.assertTrue(torch.equal(out[key], walk[key]), key) + batch = self.wrapped.get_batch(self.data) + expected = project_slots( + self.model.embedding_group, + compiled.prompt_plan, + self.model._slot_projections, + batch, + int(self.model.lm.config.hidden_size), + ) + self.assertTrue(torch.allclose(out[SLOT_EMBEDS], expected)) + self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (2, 32)) + + def test_front_end_shares_the_checkpoint_names_and_not_the_lm(self) -> None: + names = set(self.wrapped.state_dict()) + self.assertTrue(any(n.startswith("model.embedding_group.") for n in names)) + self.assertTrue(any(n.startswith("model.projections.") for n in names)) + self.assertFalse(any(".lm." in n for n in names)) + + def test_front_end_traces_and_scripts(self) -> None: + """The assembler is an FX leaf, so the export's trace-then-script works.""" + eager = self.wrapped(self.data) + scripted = torch.jit.script(symbolic_trace(self.wrapped)) + out = scripted(self.data) + for key, value in eager.items(): + self.assertTrue(torch.allclose(out[key].float(), value.float()), key) + + def test_export_assets_write_the_engine_side(self) -> None: + ckpt = write_genrec_checkpoint(self.model, os.path.join(self.test_dir, "ckpt")) + pipeline = EasyRecConfig() + pipeline.prompt_config.CopyFrom(self.prompt_config) + pipeline.data_config.label_fields.append("answer") + out_dir = os.path.join(self.test_dir, "export") + os.makedirs(out_dir) + shutil.copy(os.path.join(ckpt, "config.json"), out_dir) + GenrecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) + with open(os.path.join(out_dir, "config.json"), "r") as f: + self.assertEqual(json.load(f)["model_type"], "prompt_genrec") + for name in ( + "model.safetensors", + "prompt/prompt.json", + "prompt/tokenizer/tokenizer.json", + ): + self.assertTrue(os.path.exists(os.path.join(out_dir, name)), name) + + pipeline.export_config.use_dense_ema = True + with self.assertRaisesRegex(ValueError, "Dense EMA"): + GenrecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) + + if __name__ == "__main__": unittest.main() diff --git a/tzrec/models/model.py b/tzrec/models/model.py index 044058c0..b9eed9ef 100644 --- a/tzrec/models/model.py +++ b/tzrec/models/model.py @@ -30,6 +30,7 @@ from tzrec.features.feature import BaseFeature from tzrec.loss.pe_mtl_loss import ParetoEfficientMultiTaskLoss from tzrec.modules.utils import BaseModule +from tzrec.prompt.assembler import OUTPUT_KEYS, PROMPT_INFO_PREFIX, PromptAssembler from tzrec.protos.loss_pb2 import LossConfig from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.utils import config_util @@ -387,7 +388,12 @@ def forward( class ScriptWrapper(BaseModule): - """Model inference wrapper for jit.script.""" + """Model inference wrapper for jit.script. + + A module that exposes ``compiled_prompt`` gets its prompt assembled here, + from the parsed dict, exactly as the training collator does: one walk, two + call sites. + """ def __init__(self, module: nn.Module) -> None: super().__init__() @@ -398,6 +404,17 @@ def __init__(self, module: nn.Module) -> None: if hasattr(module, "sampler_type") else None, ) + prompt = getattr(module, "compiled_prompt", None) + self._prompt_assembler = ( + PromptAssembler( + prompt.prompt_plan, + prompt.sid_space, + plan_hash=prompt.plan_hash, + include_response=False, + ) + if prompt is not None + else None + ) @property def features(self) -> List[BaseFeature]: @@ -417,6 +434,11 @@ def get_batch( ) -> Batch: """Get batch.""" batch = self._data_parser.to_batch(data) + if self._prompt_assembler is not None: + streams = self._prompt_assembler(data) + batch.additional_infos.update( + {PROMPT_INFO_PREFIX + k: streams[k] for k in OUTPUT_KEYS} + ) batch = batch.to(device, non_blocking=True) return batch diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index a8c0cd05..422952d0 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -50,12 +50,25 @@ HOLE_SLOT_COUNTS = "hole_slot_counts" RESPONSE_LENGTHS = "response_lengths" MAX_SEQLEN = "max_seqlen" +# every stream the walk emits; a caller under FX tracing indexes these rather +# than iterating the result, which is one opaque proxy there +OUTPUT_KEYS = ( + INPUT_IDS, + CU_SEQLENS, + HOLE_POSITIONS, + HOLE_KEYS, + HOLE_SLOT_COUNTS, + RESPONSE_LENGTHS, + MAX_SEQLEN, +) # where the collator stores the streams on the batch PROMPT_INFO_PREFIX = "prompt_" PROMPT_INPUT_IDS = PROMPT_INFO_PREFIX + INPUT_IDS PROMPT_CU_SEQLENS = PROMPT_INFO_PREFIX + CU_SEQLENS PROMPT_HOLE_POSITIONS = PROMPT_INFO_PREFIX + HOLE_POSITIONS +PROMPT_HOLE_KEYS = PROMPT_INFO_PREFIX + HOLE_KEYS +PROMPT_HOLE_SLOT_COUNTS = PROMPT_INFO_PREFIX + HOLE_SLOT_COUNTS PROMPT_MAX_SEQLEN = PROMPT_INFO_PREFIX + MAX_SEQLEN PROMPT_RESPONSE_LENGTHS = PROMPT_INFO_PREFIX + RESPONSE_LENGTHS diff --git a/tzrec/prompt/export.py b/tzrec/prompt/export.py deleted file mode 100644 index 2d4027b2..00000000 --- a/tzrec/prompt/export.py +++ /dev/null @@ -1,422 +0,0 @@ -# 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. - -"""Writes what a serving runtime needs beside the HuggingFace weights. - -The HuggingFace config is rewritten into composite shape, with the backbone -nested under ``text_config``. That is not cosmetic: it names the backbone to a -runtime that composes an arbitrary causal LM behind one registered -architecture, and it is what makes an engine treat the model as carrying a -second modality, which is exactly what a projected slot is. - -The front-end is exported whether or not any slot is projected, as a standard -tzrec model directory so the online processor loads it unchanged: the scripted -module, ``fg.json``, ``pipeline.config`` and ``model_acc.json``. Lookup is a -stage of its own. By default the artifact carries the slot tables and takes -raw ids; under ``USE_DISTRIBUTED_EMBEDDING=1`` the processor's embedding stage -owns them, the module becomes the dense stage that reads its looked-up rows, -and the tables ship beside it as the sparse npz files that stage already -loads. -""" - -import copy -import glob -import json -import os -from typing import Any, Dict, List, Optional, Tuple, cast - -import torch -from torch import nn -from torchrec.modules.embedding_modules import ( - EmbeddingBagCollectionInterface, - EmbeddingCollectionInterface, -) -from torchrec.modules.mc_embedding_modules import ( - ManagedCollisionEmbeddingBagCollection, - ManagedCollisionEmbeddingCollection, -) - -from tzrec.acc import utils as acc_utils -from tzrec.features.feature import BaseFeature, create_feature_configs, create_fg_json -from tzrec.prompt.assembler import ( - CU_SEQLENS, - HOLE_KEYS, - HOLE_POSITIONS, - HOLE_SLOT_COUNTS, - INPUT_IDS, - MAX_SEQLEN, - RESPONSE_LENGTHS, - PromptAssembler, - host_lengths_key, -) -from tzrec.prompt.frontend import SLOT_EMBEDS, PromptFrontEnd, SlotTable -from tzrec.prompt.types import CompiledPrompt, FillMode, SlotSeg -from tzrec.protos.model_pb2 import FeatureGroupType -from tzrec.protos.pipeline_pb2 import EasyRecConfig -from tzrec.utils import config_util -from tzrec.utils.logging_util import logger - -SERVING_ARCH = "PromptGenRecForCausalLM" -SERVING_MODEL_TYPE = "prompt_genrec" -FRONTEND_DIR = "frontend" -SCRIPTED_MODEL_FILENAME = "scripted_model.pt" -LOOKUP_ARTIFACT = "artifact" -LOOKUP_HOST = "host" -_CONFIG = "config.json" -_SPARSE_DIR = "sparse" - - -def write_composite_config(export_dir: str) -> None: - """Rewrite ``config.json`` so the backbone sits under ``text_config``. - - Args: - export_dir: the HuggingFace export directory. - """ - path = os.path.join(export_dir, _CONFIG) - with open(path, "r") as f: - backbone: Dict[str, Any] = json.load(f) - if backbone.get("model_type") == SERVING_MODEL_TYPE: - return - composite = { - "architectures": [SERVING_ARCH], - "model_type": SERVING_MODEL_TYPE, - "text_config": backbone, - } - # a runtime that reads only the outer config still needs to size its cache - for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): - if key in backbone: - composite[key] = backbone[key] - with open(path, "w") as f: - json.dump(composite, f, indent=2) - logger.info( - f"wrote a composite config naming backbone " - f"{backbone.get('architectures', ['?'])[0]} under {SERVING_ARCH}." - ) - - -def _feature_tables(embedding_group: nn.Module) -> Dict[str, Tuple[torch.Tensor, str]]: - """Map every feature with a plain table to its weight and pooling mode. - - Managed-collision tables are skipped: their ids are remapped before the - lookup, which a plain ``EmbeddingBag`` cannot reproduce. - """ - collision_prefixes = [ - name + "." - for name, module in embedding_group.named_modules() - if isinstance( - module, - ( - ManagedCollisionEmbeddingBagCollection, - ManagedCollisionEmbeddingCollection, - ), - ) - ] - tables: Dict[str, Tuple[torch.Tensor, str]] = {} - for name, module in embedding_group.named_modules(): - if any(name.startswith(prefix) for prefix in collision_prefixes): - continue - if isinstance(module, EmbeddingBagCollectionInterface): - weights = module.state_dict() - for cfg in module.embedding_bag_configs(): - weight = weights[f"embedding_bags.{cfg.name}.weight"] - for feature_name in cfg.feature_names: - tables.setdefault( - feature_name, (weight, str(cfg.pooling.value).lower()) - ) - elif isinstance(module, EmbeddingCollectionInterface): - # a multi-value item pools its codes with the feature's own - # pooling; a single-value item is that pooling over one row - weights = module.state_dict() - for cfg in module.embedding_configs(): - weight = weights[f"embeddings.{cfg.name}.weight"] - for feature_name in cfg.feature_names: - tables.setdefault(feature_name, (weight, "")) - return tables - - -def _is_multi_valued(feature: BaseFeature) -> bool: - """Whether the parser emits ``key_lengths`` for this feature.""" - return feature.is_sequence and feature.value_dim != 1 - - -def build_slot_tables( - model: nn.Module, compiled_prompt: CompiledPrompt, features: List[BaseFeature] -) -> List[SlotTable]: - """One ``SlotTable`` per projected slot, filled from the restored model. - - Args: - model: the restored genrec model, whose ``embedding_group`` holds the - trained tables. - compiled_prompt: the compiled prompt. - features: every created feature. - - Returns: - The tables in ``projected_slots`` order. - - Raises: - ValueError: a member has no plain table the artifact could carry. - """ - tables = _feature_tables(model.embedding_group) - by_name = {feature.name: feature for feature in features} - result: List[SlotTable] = [] - for seg in compiled_prompt.prompt_plan.projected_slots: - rows: List[int] = [] - dims: List[int] = [] - modes: List[str] = [] - key_length_keys: List[str] = [] - for name in seg.feature_names: - feature = by_name[name] - if not feature.is_sparse or name not in tables: - raise ValueError( - f"prompt slot [{seg.name}] member [{name}] has no plain " - f"embedding table the front-end could carry (dense, " - f"managed-collision and dynamic tables cannot be). Export " - f"with USE_DISTRIBUTED_EMBEDDING=1 so the serving host " - f"owns the lookup." - ) - weight, mode = tables[name] - rows.append(int(weight.shape[0])) - dims.append(int(weight.shape[1])) - modes.append(mode or str(feature.pooling_type.value).lower()) - key_length_keys.append( - f"{name}.key_lengths" if _is_multi_valued(feature) else "" - ) - table = SlotTable( - [f"{name}.values" for name in seg.feature_names], - [f"{name}.lengths" for name in seg.feature_names], - key_length_keys, - rows, - dims, - seg.group_type == FeatureGroupType.JAGGED_SEQUENCE, - modes, - ) - with torch.no_grad(): - for module, name in zip(table.tables, seg.feature_names): - bag = cast(nn.EmbeddingBag, module) - bag.weight.copy_(tables[name][0].detach().to(bag.weight.dtype)) - result.append(table) - return result - - -def host_embed_keys( - compiled_prompt: CompiledPrompt, -) -> Tuple[List[List[str]], Dict[str, List[str]]]: - """Batch keys a host lookup stage fills, and the ``dense_meta`` naming them. - - The names follow the dense graph a distributed-embedding export produces: - a sequence member arrives as ``{feature}`` rows with ``{feature}__lengths``, - and a pooled slot as one ``{slot}__ebc`` tensor concatenating its members. - - Returns: - Per projected slot, the keys to concatenate on the feature axis, and - the ``dense_meta.json`` the processor reads to produce them. - """ - embed_keys: List[List[str]] = [] - dense_meta: Dict[str, List[str]] = {"sequence__ec": []} - for seg in compiled_prompt.prompt_plan.projected_slots: - if seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: - embed_keys.append(list(seg.feature_names)) - for name in seg.feature_names: - dense_meta["sequence__ec"].extend( - [f"{name}__ec", host_lengths_key(name)] - ) - else: - key = f"{seg.name}__ebc" - embed_keys.append([key]) - dense_meta[key] = [f"{name}__ebc" for name in seg.feature_names] - return embed_keys, dense_meta - - -def _frontend_inputs( - compiled_prompt: CompiledPrompt, - features: List[BaseFeature], - lookup: str, - embed_keys: List[List[str]], -) -> List[str]: - """Every batch key the exported module reads, for the serving contract.""" - by_name = {feature.name: feature for feature in features} - inputs: List[str] = [] - - def raw(name: str) -> None: - inputs.append(f"{name}.values") - inputs.append(f"{name}.lengths") - if _is_multi_valued(by_name[name]): - inputs.append(f"{name}.key_lengths") - - plan = compiled_prompt.prompt_plan - for seg in plan.segments: - if isinstance(seg, SlotSeg) and seg.fill is FillMode.INLINE: - raw(seg.feature_names[0]) - for seg, keys in zip(plan.projected_slots, embed_keys): - for name in seg.feature_names: - raw(name) - if lookup == LOOKUP_HOST: - inputs.extend(keys) - if seg.group_type == FeatureGroupType.JAGGED_SEQUENCE: - inputs.extend(host_lengths_key(name) for name in seg.feature_names) - return list(dict.fromkeys(inputs)) - - -def build_front_end( - model: nn.Module, - compiled_prompt: CompiledPrompt, - features: List[BaseFeature], - carry_tables: bool, -) -> Tuple[PromptFrontEnd, Dict[str, Any], Optional[Dict[str, List[str]]]]: - """Assemble the serving module from a restored model's parts. - - The projections are the trained modules themselves, by reference, so the - artifact carries the weights the model learned rather than a copy that can - drift from them. - - Args: - model: the restored genrec model. - compiled_prompt: the compiled prompt. - features: every created feature. - carry_tables: whether the artifact holds the slot tables, or expects a - host lookup stage to hand it rows. - - Returns: - The front-end, the ``frontend`` block of the serving contract, and the - ``dense_meta`` a host lookup stage needs (None when tables are carried). - """ - plan = compiled_prompt.prompt_plan - sid_space = compiled_prompt.sid_space - assembler = PromptAssembler( - plan, - sid_space, - plan_hash=compiled_prompt.plan_hash, - # serving has no answer to assemble; the LM generates it - include_response=False, - ) - projections = list(model._slot_projections) - if carry_tables: - lookup = LOOKUP_ARTIFACT - tables: Optional[List[SlotTable]] = build_slot_tables( - model, compiled_prompt, features - ) - embed_keys: List[List[str]] = [[] for _ in plan.projected_slots] - dense_meta: Optional[Dict[str, List[str]]] = None - else: - lookup = LOOKUP_HOST - tables = None - embed_keys, dense_meta = host_embed_keys(compiled_prompt) - front_end = PromptFrontEnd( - assembler, - projections, - embed_keys, - tables=tables, - vocab_hash=compiled_prompt.vocab_hash, - plan_hash=compiled_prompt.plan_hash, - bundle_uuid=sid_space.bundle_uuid if sid_space is not None else "", - ) - meta = { - "dir": FRONTEND_DIR, - "model": SCRIPTED_MODEL_FILENAME, - "lookup": lookup, - "inputs": _frontend_inputs(compiled_prompt, features, lookup, embed_keys), - "outputs": [ - INPUT_IDS, - CU_SEQLENS, - HOLE_POSITIONS, - HOLE_KEYS, - HOLE_SLOT_COUNTS, - SLOT_EMBEDS, - RESPONSE_LENGTHS, - MAX_SEQLEN, - ], - } - return front_end, meta, dense_meta - - -def write_front_end_dir( - front_end: PromptFrontEnd, - pipeline_config: EasyRecConfig, - features: List[BaseFeature], - export_dir: str, - dense_meta: Optional[Dict[str, List[str]]] = None, -) -> str: - """Script the front-end into ``frontend/``, laid out as a tzrec model dir. - - Args: - front_end: the module to export. - pipeline_config: the pipeline config, whose feature configs are - rewritten beside the copied assets. - features: every created feature. - export_dir: the HuggingFace export directory. - dense_meta: the processor's dense-stage input map, written when a host - lookup stage feeds the module. - - Returns: - The ``frontend/`` directory written. - """ - frontend_dir = os.path.join(export_dir, FRONTEND_DIR) - os.makedirs(frontend_dir, exist_ok=True) - torch.jit.script(front_end.eval()).save( - os.path.join(frontend_dir, SCRIPTED_MODEL_FILENAME) - ) - - feature_configs = create_feature_configs(features, asset_dir=frontend_dir) - served_config = copy.deepcopy(pipeline_config) - served_config.ClearField("feature_configs") - served_config.feature_configs.extend(feature_configs) - config_util.save_message( - served_config, os.path.join(frontend_dir, "pipeline.config") - ) - with open(os.path.join(frontend_dir, "fg.json"), "w") as f: - json.dump(create_fg_json(features, asset_dir=frontend_dir), f, indent=4) - with open(os.path.join(frontend_dir, "model_acc.json"), "w") as f: - json.dump(acc_utils.export_acc_config(), f, indent=4) - if dense_meta is not None: - with open(os.path.join(frontend_dir, "dense_meta.json"), "w") as f: - json.dump(dense_meta, f, indent=4) - logger.info(f"wrote the scripted prompt front-end to {frontend_dir}.") - return frontend_dir - - -def export_sparse_tables( - wrapped_model: nn.Module, checkpoint_path: str, frontend_dir: str -) -> None: - """Write the slot tables as the sparse npz files a host lookup stage loads. - - Single rank, in the layout ``export_distributed_embedding`` produces. - - Args: - wrapped_model: the restored model under its inference wrapper, so the - table names match the checkpoint's. - checkpoint_path: the checkpoint, read for dynamic tables. - frontend_dir: the ``frontend/`` directory to write ``sparse/`` into. - """ - from tzrec.utils import export_util, npz_util - - bag_info, emb_info = export_util._get_sparse_table_to_embedding_info(wrapped_model) - local, dynamic, emb_meta, feat_meta = export_util._get_sparse_embedding_tensor( - wrapped_model, checkpoint_path, emb_info, bag_info - ) - sparse_dir = os.path.join(frontend_dir, _SPARSE_DIR) - os.makedirs(sparse_dir, exist_ok=True) - shard = "sparse_embeddings-00-of-01" - npz_util.savez_streaming(os.path.join(sparse_dir, f"{shard}.npz"), local) - if dynamic: - npz_util.savez_streaming( - os.path.join(sparse_dir, "sparse_dynamic_embedding-00-of-01.npz"), dynamic - ) - with open(os.path.join(sparse_dir, f"{shard}.json"), "w") as f: - json.dump(emb_meta, f, indent=4) - with open(os.path.join(sparse_dir, "sparse_features.json"), "w") as f: - json.dump(feat_meta, f, indent=4) - shards = [] - for path in glob.glob(os.path.join(sparse_dir, "sparse_embeddings*.json")): - with open(path, "r") as f: - shards.append(json.load(f)) - with open(os.path.join(sparse_dir, "sparse_embedding.json"), "w") as f: - json.dump(export_util._merge_sharded_embedding_json(shards), f, indent=4) - logger.info(f"wrote {len(local)} slot table(s) to {sparse_dir}.") diff --git a/tzrec/prompt/export_test.py b/tzrec/prompt/export_test.py deleted file mode 100644 index 204caaef..00000000 --- a/tzrec/prompt/export_test.py +++ /dev/null @@ -1,249 +0,0 @@ -# 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 json -import os -import unittest - -import torch -from torchrec import KeyedJaggedTensor - -from tzrec.datasets.utils import BASE_DATA_GROUP, Batch -from tzrec.prompt.export import ( - LOOKUP_ARTIFACT, - LOOKUP_HOST, - build_front_end, - build_slot_tables, - host_embed_keys, - write_composite_config, - write_front_end_dir, -) -from tzrec.prompt.persist import write_serving_contract -from tzrec.protos.pipeline_pb2 import EasyRecConfig -from tzrec.tests.prompt_test_util import ( - _HIST, - GenrecModelTestBase, - create_prompt_feature, - offset_sid_codes, - projected_feature, -) -from tzrec.utils.state_dict_util import init_parameters - -_CAT = ( - 'id_feature { feature_name: "cat" expression: "user:cat" num_buckets: 16 ' - 'embedding_dim: 8 pooling: "mean" }' -) -_VEC = 'raw_feature { feature_name: "vec" expression: "user:vec" value_dim: 4 }' -_TEMPLATE = "History : {{hist}} . {{beh}} {{cat}} Predict :" - - -class PromptExportTest(GenrecModelTestBase): - """Builds the serving pieces from an in-memory model with projected slots.""" - - def setUp(self) -> None: - super().setUp() - self.features = [ - create_prompt_feature(_HIST), - create_prompt_feature(projected_feature("beh", 8)), - create_prompt_feature(_CAT), - ] - self.compiled_prompt = self._compile(self.features, template=_TEMPLATE) - self.model = self._model() - init_parameters(self.model, torch.device("cpu")) - self.raw = { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2, 3, 0, 1], [4, 4, 4]) - ), - "hist.lengths": torch.tensor([6]), - "beh.values": torch.tensor([3, 9]), - "beh.lengths": torch.tensor([2]), - "cat.values": torch.tensor([5]), - "cat.lengths": torch.tensor([1]), - } - sparse = KeyedJaggedTensor( - keys=["beh", "cat"], - values=torch.tensor([3, 9, 5]), - lengths=torch.tensor([2, 1]), - ) - self.grouped = self.model.embedding_group( - Batch(sparse_features={BASE_DATA_GROUP: sparse}) - ) - - def _expected_slot_embeds(self) -> torch.Tensor: - projections = self.model._slot_projections - return torch.cat( - [ - projections[0](self.grouped["beh.sequence"]), - projections[1](self.grouped["cat"]), - ] - ) - - def test_slot_tables_reproduce_the_embedding_group(self) -> None: - """The carried tables look up exactly what the training model did.""" - tables = build_slot_tables(self.model, self.compiled_prompt, self.features) - self.assertEqual(len(tables), 2) - self.assertTrue( - torch.allclose(tables[0](self.raw), self.grouped["beh.sequence"]) - ) - self.assertTrue(torch.allclose(tables[1](self.raw), self.grouped["cat"])) - - def test_the_artifact_shape_matches_the_model_on_the_same_batch(self) -> None: - front_end, meta, dense_meta = build_front_end( - self.model, self.compiled_prompt, self.features, carry_tables=True - ) - out = front_end(self.raw) - self.assertTrue( - torch.allclose(out["slot_embeds"], self._expected_slot_embeds(), atol=1e-6) - ) - self.assertEqual(out["hole_slot_counts"].tolist(), [2, 1]) - self.assertEqual(meta["lookup"], LOOKUP_ARTIFACT) - self.assertIsNone(dense_meta) - self.assertEqual( - meta["inputs"], - [ - "hist.values", - "hist.lengths", - "beh.values", - "beh.lengths", - "cat.values", - "cat.lengths", - ], - ) - self.assertEqual(front_end.vocab_hash, self.compiled_prompt.vocab_hash) - - def test_the_host_shape_reads_the_processor_keys(self) -> None: - """A host lookup stage hands over rows named as dense_meta describes.""" - artifact, _, _ = build_front_end( - self.model, self.compiled_prompt, self.features, carry_tables=True - ) - host, meta, dense_meta = build_front_end( - self.model, self.compiled_prompt, self.features, carry_tables=False - ) - self.assertEqual(meta["lookup"], LOOKUP_HOST) - self.assertEqual( - dense_meta, - {"sequence__ec": ["beh__ec", "beh__lengths"], "cat__ebc": ["cat__ebc"]}, - ) - self.assertEqual( - host_embed_keys(self.compiled_prompt)[0], [["beh"], ["cat__ebc"]] - ) - batch = dict(self.raw) - batch["beh"] = self.grouped["beh.sequence"] - batch["beh__lengths"] = self.raw["beh.lengths"] - batch["cat__ebc"] = self.grouped["cat"] - expected = artifact(self.raw) - got = torch.jit.script(host.eval())(batch) - self.assertTrue(torch.allclose(got["slot_embeds"], expected["slot_embeds"])) - self.assertTrue(torch.equal(got["hole_keys"], expected["hole_keys"])) - self.assertIn("beh__lengths", meta["inputs"]) - self.assertIn("cat__ebc", meta["inputs"]) - - def test_a_member_without_a_plain_table_cannot_be_carried(self) -> None: - features = [create_prompt_feature(_HIST), create_prompt_feature(_VEC)] - compiled = self._compile(features, template="History : {{hist}} {{vec}} :") - model = self._model(features=features, compiled_prompt=compiled) - init_parameters(model, torch.device("cpu")) - with self.assertRaisesRegex(ValueError, "USE_DISTRIBUTED_EMBEDDING=1"): - build_slot_tables(model, compiled, features) - - def test_write_composite_config_is_idempotent(self) -> None: - path = os.path.join(self.test_dir, "config.json") - backbone = { - "model_type": "qwen2", - "architectures": ["Qwen2ForCausalLM"], - "hidden_size": 32, - "vocab_size": 77, - } - with open(path, "w") as f: - json.dump(backbone, f) - write_composite_config(self.test_dir) - write_composite_config(self.test_dir) - with open(path, "r") as f: - composite = json.load(f) - self.assertEqual(composite["architectures"], ["PromptGenRecForCausalLM"]) - self.assertEqual(composite["model_type"], "prompt_genrec") - self.assertEqual(composite["text_config"], backbone) - self.assertEqual(composite["hidden_size"], 32) - - def test_write_front_end_dir_lays_out_a_tzrec_model_dir(self) -> None: - front_end, _, _ = build_front_end( - self.model, self.compiled_prompt, self.features, carry_tables=True - ) - frontend_dir = write_front_end_dir( - front_end, - EasyRecConfig(model_dir="unused"), - self.features, - self.test_dir, - dense_meta={"sequence__ec": ["beh__ec", "beh__lengths"]}, - ) - for name in ( - "scripted_model.pt", - "fg.json", - "pipeline.config", - "model_acc.json", - "dense_meta.json", - ): - self.assertTrue(os.path.exists(os.path.join(frontend_dir, name)), name) - with open(os.path.join(frontend_dir, "model_acc.json"), "r") as f: - self.assertEqual(json.load(f)["SPARSE_INT64"], "1") - with open(os.path.join(frontend_dir, "fg.json"), "r") as f: - names = [feature["feature_name"] for feature in json.load(f)["features"]] - self.assertEqual(names, ["hist", "beh", "cat"]) - loaded = torch.jit.load(os.path.join(frontend_dir, "scripted_model.pt")) - out = loaded(self.raw, torch.device("cpu")) - self.assertTrue( - torch.allclose(out["slot_embeds"], self._expected_slot_embeds(), atol=1e-6) - ) - - def test_serving_contract_lists_what_the_runtime_reads(self) -> None: - _, meta, _ = build_front_end( - self.model, self.compiled_prompt, self.features, carry_tables=True - ) - path = write_serving_contract(self.compiled_prompt, self.test_dir, meta) - with open(path, "r") as f: - contract = json.load(f) - space = self.compiled_prompt.sid_space - self.assertEqual(contract["sid_space"]["band_lo"], list(space.band_lo)) - self.assertEqual( - contract["sid_space"]["base_vocab_size"], space.base_vocab_size - ) - self.assertEqual(contract["sid_space"]["bundle_uuid"], "") - self.assertEqual( - contract["decode"], - [ - { - "level": level, - "band_lo": space.band_lo[level], - "band_hi": space.band_hi[level], - } - for level in range(3) - ], - ) - self.assertEqual(contract["vocab_hash"], self.compiled_prompt.vocab_hash) - self.assertEqual(contract["plan_hash"], self.compiled_prompt.plan_hash) - self.assertEqual(contract["frontend"], meta) - self.assertEqual(contract["static_prefix_len"], 2) - - def test_a_static_run_in_the_response_is_a_forced_token(self) -> None: - compiled = self._compile( - self.features, template=_TEMPLATE, response="Predict : {{answer}}" - ) - path = write_serving_contract(compiled, self.test_dir, {}) - with open(path, "r") as f: - decode = json.load(f)["decode"] - self.assertEqual(len(decode), 5) - self.assertEqual(decode[0], {"token_id": 1}) - self.assertEqual(decode[1], {"token_id": 2}) - self.assertIn("band_lo", decode[2]) - - -if __name__ == "__main__": - unittest.main() diff --git a/tzrec/prompt/frontend.py b/tzrec/prompt/frontend.py deleted file mode 100644 index 4d0852e9..00000000 --- a/tzrec/prompt/frontend.py +++ /dev/null @@ -1,215 +0,0 @@ -# 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. - -"""The prompt front-end a serving runtime loads: assemble, look up, project. - -It wraps the collator's own ``PromptAssembler`` -- the one walk, scripted -- -and adds the two stages serving needs after it: the slot lookup, when the host -has no embedding stage of its own, and the trained projections into the LM -input space. -""" - -from typing import Dict, List, Optional - -import torch -from torch import nn - -from tzrec.prompt.assembler import PromptAssembler, batch_device - -SLOT_EMBEDS = "slot_embeds" - - -class SlotTable(nn.Module): - """One projected slot's members, as plain embedding tables. - - Used when the host does not run an embedding stage of its own. The lookup - has to live somewhere: either the host does it and hands over vectors, or - the artifact carries the tables. Keeping both shapes behind one module means - the walk, the fold and the projections do not change between them. - - Args: - value_keys: each member's value key. - length_keys: each member's per-row counts. - key_length_keys: each member's per-item value counts, empty when the - member holds one value per item. - num_embeddings: each member's table size. - dims: each member's embedding dimension. - is_sequence: whether holes are items rather than rows. - modes: each member's pooling, ``sum`` or ``mean``. - """ - - value_keys: List[str] - length_keys: List[str] - key_length_keys: List[str] - - def __init__( - self, - value_keys: List[str], - length_keys: List[str], - key_length_keys: List[str], - num_embeddings: List[int], - dims: List[int], - is_sequence: bool, - modes: Optional[List[str]] = None, - ) -> None: - super().__init__() - self.value_keys = value_keys - self.length_keys = length_keys - self.key_length_keys = key_length_keys - self.is_sequence = is_sequence - modes = modes if modes is not None else ["sum"] * len(dims) - self.tables = nn.ModuleList( - [ - nn.EmbeddingBag(rows, dim, mode=mode, include_last_offset=True) - for rows, dim, mode in zip(num_embeddings, dims, modes) - ] - ) - - def forward(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor: - """Look every member up and concatenate along the feature axis. - - Returns: - ``(num_holes, group_total_dim)``, ordered by hole. - """ - parts: List[torch.Tensor] = [] - index = 0 - for table in self.tables: - values = batch[self.value_keys[index]].to(torch.int64).reshape(-1) - key = self.key_length_keys[index] - if key != "": - counts = batch[key].to(torch.int64) - elif self.is_sequence: - counts = torch.ones_like(values) - else: - counts = batch[self.length_keys[index]].to(torch.int64) - offsets = torch.cat( - [ - torch.zeros(1, dtype=torch.int64, device=values.device), - torch.cumsum(counts, dim=0), - ] - ) - parts.append(table(values, offsets)) - index += 1 - return torch.cat(parts, dim=1) - - -class PromptFrontEnd(nn.Module): - """The serving artifact: assemble, look up, project. - - Exported whether or not any slot is projected -- with none it degenerates - to the assembler and three empty streams, which is cheap and keeps what a - prompt happens to contain from deciding whether a serving-critical artifact - exists. - - ``slot_embeds`` is ordered by ascending hole, matching ``hole_positions`` - entry for entry. That ordering is an obligation rather than a check: the - engine scatters positionally, so a permuted source gives every hole a - neighbour's embedding, the counts still agree, and nothing raises. - - Lookup is a stage of its own. Without a host embedding stage the artifact - carries the tables and the batch holds raw ids. Under a host stage the - embeddings are already in the batch when this runs, one entry per key in - ``embed_keys`` for each slot, concatenated on the feature axis. Everything - downstream of the lookup is identical either way. - - Args: - assembler: the walk. - projections: one per projected slot in ``projected_slots`` order; slots - sharing a module appear more than once, by reference. - embed_keys: per projected slot, the batch keys holding its members' - looked-up rows, in member order. Ignored when ``tables`` is given. - tables: one per projected slot, when the artifact carries the slot - tables rather than receiving vectors from a host stage. - vocab_hash: the compiled prompt's vocabulary digest, so a loader can - refuse a front-end that does not pair with its ``prompt.json``. - plan_hash: the compiled prompt's plan digest. - bundle_uuid: identity of the SID bundle the prompt was compiled against. - """ - - embed_keys: List[List[str]] - vocab_hash: str - plan_hash: str - bundle_uuid: str - - def __init__( - self, - assembler: PromptAssembler, - projections: List[nn.Module], - embed_keys: List[List[str]], - tables: Optional[List[SlotTable]] = None, - vocab_hash: str = "", - plan_hash: str = "", - bundle_uuid: str = "", - ) -> None: - super().__init__() - self.assembler = assembler - self.projections = nn.ModuleList(projections) - self.embed_keys = embed_keys - self.has_tables = tables is not None - self.tables = nn.ModuleList(tables if tables is not None else []) - self.vocab_hash = vocab_hash - self.plan_hash = plan_hash - self.bundle_uuid = bundle_uuid - - def forward( - self, - data: Dict[str, torch.Tensor], - device: Optional[torch.device] = None, - ) -> Dict[str, torch.Tensor]: - """Assemble one batch and project its holes. - - The second argument is not decoration: a C++ host that runs this as its - JIT stage calls ``forward(data, device)``, the same pair tzrec's own - ScriptWrapper takes. A Python caller may omit it, in which case the - batch stays where it is. - - Args: - data: the parsed feature dict, plus looked-up rows under a host - lookup stage. - device: where to run; the batch is moved there first. - - Returns: - The assembler's outputs plus ``slot_embeds``. - """ - target = batch_device(data) if device is None else device - batch: Dict[str, torch.Tensor] = {} - for key, value in data.items(): - batch[key] = value.to(target) - out = self.assembler(batch) - - # gathered first, into a plain list: TorchScript indexes a ModuleList - # only with a literal, so the two lists cannot be walked together - features: List[torch.Tensor] = [] - if self.has_tables: - for table in self.tables: - features.append(table(batch)) - else: - for keys in self.embed_keys: - members: List[torch.Tensor] = [] - for key in keys: - members.append(batch[key]) - features.append(torch.cat(members, dim=1)) - - parts: List[torch.Tensor] = [] - index = 0 - for projection in self.projections: - parts.append(projection(features[index])) - index += 1 - - # literals rather than the module constants: TorchScript cannot see a - # module-level global. ``frontend_test`` pins them together. - if len(parts) > 0: - out["slot_embeds"] = torch.cat(parts, dim=0) - else: - out["slot_embeds"] = torch.zeros( - 0, 0, dtype=torch.float32, device=out["input_ids"].device - ) - return out diff --git a/tzrec/prompt/frontend_test.py b/tzrec/prompt/frontend_test.py deleted file mode 100644 index a171a113..00000000 --- a/tzrec/prompt/frontend_test.py +++ /dev/null @@ -1,180 +0,0 @@ -# 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 unittest - -import numpy as np -import torch - -from tzrec.prompt.assembler import HOLE_KEYS, PromptAssembler, host_lengths_key -from tzrec.prompt.frontend import SLOT_EMBEDS, PromptFrontEnd, SlotTable -from tzrec.prompt.types import ( - FillMode, - PromptPlan, - ResolvedSidSpace, - SlotSeg, - Static, - Width, - WidthKind, -) -from tzrec.protos.model_pb2 import FeatureGroupType -from tzrec.utils.test_util import make_test_dir - -_BASE_VOCAB = 1000 -_CODEBOOK = (4, 4, 4) -_LEVEL_OFFSETS = (0, 4, 8) - - -def _sid_space() -> ResolvedSidSpace: - return ResolvedSidSpace( - codebook=_CODEBOOK, - num_levels=3, - base_vocab_size=_BASE_VOCAB, - level_offsets=_LEVEL_OFFSETS, - band_lo=tuple(_BASE_VOCAB + o for o in _LEVEL_OFFSETS), - band_hi=tuple( - _BASE_VOCAB + o + c - 1 for o, c in zip(_LEVEL_OFFSETS, _CODEBOOK) - ), - target_vocab_size=_BASE_VOCAB + 13, - sentinel_token_id=_BASE_VOCAB + 12, - eos_token_id=2, - pad_token_id=3, - bundle_uuid="test-bundle", - ) - - -def _slot(slot_id, name, fill): - return SlotSeg( - slot_id=slot_id, - name=name, - feature_names=(name,), - group_type=FeatureGroupType.JAGGED_SEQUENCE, - output_key=".sequence", - fill=fill, - width=Width(WidthKind.BOUNDED, 30), - ) - - -def _plan(segments, projected=()): - return PromptPlan( - segments=segments, - response_segments=(), - max_length=256, - max_total_length=None, - max_holes=30, - logits_suffix_len=4, - static_prefix_len=1, - projected_slots=projected, - ) - - -def _tensors(raw): - return {key: torch.from_numpy(np.asarray(value)) for key, value in raw.items()} - - -class FrontEndTest(unittest.TestCase): - def setUp(self): - self.slot = _slot(0, "beh", FillMode.PROJECTED) - self.plan = _plan(segments=(Static((10,)), self.slot), projected=(self.slot,)) - self.batch = _tensors( - { - "beh.values": np.array([1, 2, 3], dtype=np.int64), - "beh.lengths": np.array([3], dtype=np.int64), - } - ) - - def _table(self): - torch.manual_seed(0) - return SlotTable(["beh.values"], ["beh.lengths"], [""], [16], [8], True) - - def _front_end(self, tables, embed_keys=(("beh",),)): - torch.manual_seed(1) - assembler = PromptAssembler(self.plan, _sid_space(), plan_hash="a1b2") - projection = torch.nn.Linear(8, 16) - return PromptFrontEnd( - assembler, - [projection], - [list(keys) for keys in embed_keys], - tables=tables, - vocab_hash="vocab", - plan_hash="plan", - bundle_uuid="test-bundle", - ) - - def test_in_module_tables_produce_one_embedding_per_hole(self): - """Without a host lookup stage, the artifact carries the tables.""" - out = self._front_end([self._table()])(self.batch) - self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (3, 16)) - self.assertEqual(int(out["hole_positions"].numel()), 3) - - def test_both_lookup_shapes_agree_given_the_same_rows(self): - """Where the lookup happens must not change what the model sees.""" - table = self._table() - with_tables = self._front_end([table]) - host = self._front_end(None) - # a host stage hands over rows and item counts, not ids and .lengths - batch = { - "beh.values": self.batch["beh.values"], - "beh": table(self.batch), - host_lengths_key("beh"): self.batch["beh.lengths"], - } - self.assertTrue( - torch.allclose( - with_tables(self.batch)[SLOT_EMBEDS], host(batch)[SLOT_EMBEDS] - ) - ) - self.assertTrue( - torch.equal(with_tables(self.batch)[HOLE_KEYS], host(batch)[HOLE_KEYS]) - ) - - def test_the_device_argument_is_optional(self): - """The processor passes a device like ScriptWrapper; sglang does not.""" - module = torch.jit.script(self._front_end([self._table()]).eval()) - implicit = module(self.batch) - explicit = module(self.batch, torch.device("cpu")) - for key, value in implicit.items(): - self.assertTrue(torch.equal(explicit[key], value), key) - - def test_the_artifact_reloads_with_its_identity(self): - """A loader pairs the front-end with prompt.json by these attributes.""" - path = os.path.join(make_test_dir(), "scripted_model.pt") - torch.jit.script(self._front_end([self._table()]).eval()).save(path) - loaded = torch.jit.load(path) - self.assertEqual( - (loaded.vocab_hash, loaded.plan_hash, loaded.bundle_uuid), - ("vocab", "plan", "test-bundle"), - ) - out = loaded(self.batch) - self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (3, 16)) - - def test_pattern_i_exports_empty_hole_streams(self): - """With no projected slot the artifact degenerates, and still exists.""" - hist = _slot(0, "hist", FillMode.INLINE) - plan = _plan(segments=(Static((10,)), hist)) - module = torch.jit.script( - PromptFrontEnd(PromptAssembler(plan, _sid_space()), [], []).eval() - ) - out = module( - _tensors( - { - "hist.values": np.array([0, 5, 10], dtype=np.int64), - "hist.lengths": np.array([3], dtype=np.int64), - } - ) - ) - self.assertEqual(out["input_ids"].tolist(), [10, 1000, 1005, 1010]) - self.assertEqual(int(out["hole_positions"].numel()), 0) - self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (0, 0)) - - -if __name__ == "__main__": - unittest.main() diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index 5d74c52d..2a877b92 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -29,7 +29,7 @@ from tzrec.prompt.assembler import mix64 from tzrec.prompt.types import FoldConstants -from tzrec.tests.prompt_test_util import export_tiny_genrec, offset_sid_codes +from tzrec.tests.prompt_test_util import export_tiny_genrec from tzrec.utils.test_util import make_test_dir # sglang's multimodal/processors/prompt_genrec.py folds a slot's per-hole keys @@ -63,20 +63,14 @@ def setUpClass(cls) -> None: ) as f: cls.contract = json.load(f) cls.front_end = torch.jit.load( - os.path.join(cls.exported.export_dir, "frontend", "scripted_model.pt") + os.path.join(cls.exported.export_dir, "scripted_model.pt") ) def _payload(self): """A request as the in-process processor receives it: an npz blob.""" buffer = io.BytesIO() np.savez_compressed( - buffer, - **{ - "hist.values": offset_sid_codes([0, 1, 2, 3, 0, 1], [4, 4, 4]), - "hist.lengths": np.array([6], dtype=np.int64), - "beh.values": np.array([3, 9], dtype=np.int64), - "beh.lengths": np.array([2], dtype=np.int64), - }, + buffer, **{k: v.numpy() for k, v in self.exported.sample_data.items()} ) buffer.seek(0) with np.load(buffer, allow_pickle=False) as data: @@ -107,7 +101,10 @@ def test_weights_keep_the_backbones_own_names(self) -> None: self.assertIn("model.embed_tokens.weight", exported) def test_front_end_output_feeds_the_processor(self) -> None: + # sglang calls it with the batch alone; the processor adds a device out = self.front_end(self._payload()) + with_device = self.front_end(self._payload(), torch.device("cpu")) + self.assertTrue(torch.equal(with_device["hole_keys"], out["hole_keys"])) for key in ( "input_ids", "hole_positions", @@ -152,12 +149,12 @@ def test_prompt_json_carries_the_index_builder_inputs(self) -> None: self.assertEqual( self.contract["vocab_hash"], self.exported.compiled_prompt.vocab_hash ) - self.assertEqual(self.contract["frontend"]["dir"], "frontend") self.assertEqual(self.contract["frontend"]["model"], "scripted_model.pt") - # the artifact pairs with the contract by identity, not by path - self.assertEqual(self.front_end.vocab_hash, self.contract["vocab_hash"]) - self.assertEqual(self.front_end.plan_hash, self.contract["plan_hash"]) - self.assertEqual(self.front_end.bundle_uuid, "bundle-test") + self.assertEqual(self.contract["frontend"]["lookup"], "artifact") + self.assertIn("beh.values", self.contract["frontend"]["inputs"]) + self.assertEqual( + self.contract["plan_hash"], self.exported.compiled_prompt.plan_hash + ) def test_tokenizer_dir_decodes_a_sid_atom(self) -> None: from transformers import AutoTokenizer diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index 7aa3046f..430dc13b 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -11,15 +11,15 @@ import json import os +import subprocess +import sys import unittest -from unittest import mock import numpy as np import torch import torch.fx from tzrec.datasets.utils import Batch -from tzrec.main import export from tzrec.models.model import TrainWrapper from tzrec.prompt.assembler import ( CU_SEQLENS, @@ -34,7 +34,7 @@ export_tiny_genrec, offset_sid_codes, ) -from tzrec.utils.test_util import make_test_dir +from tzrec.utils.test_util import gpu_unavailable, make_test_dir, mark_ci_scope _CODEBOOK = [4, 4, 4] _WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] @@ -89,101 +89,130 @@ def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(model)) -def _serving_batch(): - """One request as the front-end reads it: an INLINE history, one behaviour.""" - return { - "hist.values": torch.tensor(offset_sid_codes([0, 1, 2, 3, 0, 1], _CODEBOOK)), - "hist.lengths": torch.tensor([6]), - "beh.values": torch.tensor([3, 9]), - "beh.lengths": torch.tensor([2]), - } - - class GenrecExportIntegrationTest(unittest.TestCase): - """checkpoint -> export -> the artifacts a serving runtime loads.""" + """checkpoint -> export -> the artifacts an LLM engine and a processor load.""" def setUp(self) -> None: self.test_dir = make_test_dir() - def test_export_writes_a_loadable_serving_directory(self) -> None: - exported = export_tiny_genrec(self.test_dir) + def test_export_writes_one_serving_directory(self) -> None: + exported = export_tiny_genrec(self.test_dir, env={"QUANT_EMB": "0"}) for name in ( + "scripted_model.pt", + "fg.json", + "pipeline.config", + "model_acc.json", "config.json", "model.safetensors", "prompt/prompt.json", "prompt/tokenizer/tokenizer.json", "prompt/tokenizer/tokenizer_config.json", - "frontend/scripted_model.pt", - "frontend/fg.json", - "frontend/pipeline.config", - "frontend/model_acc.json", ): self.assertTrue( os.path.exists(os.path.join(exported.export_dir, name)), name ) - self.assertFalse( - os.path.exists(os.path.join(exported.export_dir, "frontend/sparse")) - ) + self.assertFalse(os.path.exists(os.path.join(exported.export_dir, "sparse"))) - # the collator and the exported artifact are two call sites of one walk - batch = _serving_batch() + # the exported artifact and the collator are two call sites of one walk + data = exported.sample_data compiled = exported.compiled_prompt collator = PromptAssembler( compiled.prompt_plan, compiled.sid_space, plan_hash=compiled.plan_hash, include_response=False, - )(batch) + )(data) front_end = torch.jit.load( - os.path.join(exported.export_dir, "frontend", "scripted_model.pt") + os.path.join(exported.export_dir, "scripted_model.pt") ) - out = front_end(batch, torch.device("cpu")) + out = front_end(data, torch.device("cpu")) for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, HOLE_KEYS): self.assertTrue(torch.equal(out[key], collator[key]), key) - self.assertEqual(tuple(out["slot_embeds"].shape), (2, 32)) + self.assertEqual( + tuple(out["slot_embeds"].shape), (int(out[HOLE_POSITIONS].numel()), 32) + ) + @unittest.skipIf(*gpu_unavailable) + @mark_ci_scope("gpu") def test_distributed_embedding_export_writes_the_processor_shape(self) -> None: - exported = export_tiny_genrec(self.test_dir) + exported = export_tiny_genrec(self.test_dir, env={"QUANT_EMB": "0"}) dist_dir = os.path.join(self.test_dir, "export_dist") - with mock.patch.dict(os.environ, {"USE_DISTRIBUTED_EMBEDDING": "1"}): - export( - exported.config_path, dist_dir, checkpoint_path=exported.checkpoint_dir + # its own process: the planner wants a fresh nccl group on the device + env = dict(os.environ) + env.update( + { + "PYTHONPATH": ".", + "USE_DISTRIBUTED_EMBEDDING": "1", + "QUANT_EMB": "0", + "MASTER_ADDR": "127.0.0.1", + "MASTER_PORT": os.environ.get("MASTER_PORT", "29511"), + "RANK": "0", + "LOCAL_RANK": "0", + "WORLD_SIZE": "1", + } + ) + log_path = os.path.join(self.test_dir, "export_dist.log") + with open(log_path, "w") as log: + code = subprocess.call( + [ + sys.executable, + "-m", + "tzrec.export", + "--pipeline_config_path", + exported.config_path, + "--export_dir", + dist_dir, + "--checkpoint_path", + exported.checkpoint_dir, + ], + env=env, + stdout=log, + stderr=subprocess.STDOUT, ) - frontend_dir = os.path.join(dist_dir, "frontend") - sparse_dir = os.path.join(frontend_dir, "sparse") + with open(log_path, "r") as log: + self.assertEqual(code, 0, log.read()[-4000:]) + sparse_dir = os.path.join(dist_dir, "sparse") for name in ( "sparse_embeddings-00-of-01.npz", - "sparse_embeddings-00-of-01.json", "sparse_embedding.json", "sparse_features.json", ): self.assertTrue(os.path.exists(os.path.join(sparse_dir, name)), name) - with open(os.path.join(frontend_dir, "model_acc.json"), "r") as f: - acc = json.load(f) - self.assertEqual(acc["DISTRIBUTED_EMBEDDING"], "1") - self.assertEqual(acc["INPUT_TILE"], "3") - with open(os.path.join(frontend_dir, "dense_meta.json"), "r") as f: + with open(os.path.join(dist_dir, "model_acc.json"), "r") as f: + self.assertEqual(json.load(f)["DISTRIBUTED_EMBEDDING"], "1") + with open(os.path.join(dist_dir, "dense_meta.json"), "r") as f: dense_meta = json.load(f) - self.assertEqual(dense_meta, {"sequence__ec": ["beh__ec", "beh__lengths"]}) + self.assertEqual(dense_meta["sequence__ec"], ["beh__ec", "beh__lengths"]) with open(os.path.join(dist_dir, "prompt", "prompt.json"), "r") as f: self.assertEqual(json.load(f)["frontend"]["lookup"], "host") - # a processor simulator: look the ids up in the exported tables the way - # the distributed-embedding stage does, then feed the dense stage - batch = _serving_batch() + # a processor simulator: one request is one user, looked up in the + # exported tables the way the distributed-embedding stage does, then + # fed to the dense stage input-tiled with one candidate + sample = exported.sample_data + request = { + "hist.values": sample["hist.values"][: int(sample["hist.lengths"][0])], + "hist.lengths": sample["hist.lengths"][:1], + "beh.values": sample["beh.values"][: int(sample["beh.lengths"][0])], + "beh.lengths": sample["beh.lengths"][:1], + } with open(os.path.join(sparse_dir, "sparse_features.json"), "r") as f: table_name = json.load(f)["beh__ec"]["embedding_name"] with np.load(os.path.join(sparse_dir, "sparse_embeddings-00-of-01.npz")) as npz: table = npz[table_name] - host = dict(batch) - host["beh"] = torch.from_numpy( - table[batch["beh.values"].numpy()].astype(np.float32) + data = dict(request) + data["beh"] = torch.from_numpy( + table[request["beh.values"].numpy()].astype(np.float32) ) - host["beh__lengths"] = batch["beh.lengths"] - got = torch.jit.load(os.path.join(frontend_dir, "scripted_model.pt"))(host) + data["beh__lengths"] = request["beh.lengths"] + data["batch_size"] = torch.tensor(1) + # the planner exported on the GPU; the processor loads onto its device + got = torch.jit.load( + os.path.join(dist_dir, "scripted_model.pt"), map_location="cpu" + )(data) expected = torch.jit.load( - os.path.join(exported.export_dir, "frontend", "scripted_model.pt") - )(batch) + os.path.join(exported.export_dir, "scripted_model.pt") + )(request) self.assertTrue( torch.allclose(got["slot_embeds"], expected["slot_embeds"], atol=1e-6) ) diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index b15697d8..d4ff622a 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -19,6 +19,7 @@ import torch from google.protobuf import text_format from tokenizers import Tokenizer, models, pre_tokenizers +from torch import distributed as dist from tzrec.datasets.utils import BASE_DATA_GROUP, Batch from tzrec.features.feature import BaseFeature, FgMode, create_features @@ -193,7 +194,7 @@ def write_genrec_checkpoint(model: torch.nn.Module, ckpt_dir: str) -> str: @dataclasses.dataclass class ExportedGenrec: - """A tiny genrec model, its checkpoint and its HF export. + """A tiny genrec model, its checkpoint and its export. Args: config: the pipeline config the export ran on. @@ -202,6 +203,8 @@ class ExportedGenrec: compiled_prompt: the prompt as the export compiled it. checkpoint_dir: the checkpoint the export converted. export_dir: the export directory. + sample_data: one parsed batch as the served module reads it, the dict + ``Batch.to_dict`` emits for the mock data. """ config: EasyRecConfig @@ -210,30 +213,72 @@ class ExportedGenrec: compiled_prompt: CompiledPrompt checkpoint_dir: str export_dir: str + sample_data: Dict[str, torch.Tensor] + + +def _write_mock_prompt_data(path: str, projected: bool, rows: int = 8) -> str: + """Write a parquet of SID histories the mock prompt reads. + + Returns: + The glob a data config points at. + """ + import pyarrow as pa + import pyarrow.parquet as pq + + rng = np.random.default_rng(0) + columns = { + "hist": [ + offset_sid_codes(rng.integers(0, 4, size=6), _CODEBOOK).tolist() + for _ in range(rows) + ], + "answer": [ + offset_sid_codes(rng.integers(0, 4, size=3), _CODEBOOK).tolist() + for _ in range(rows) + ], + } + if projected: + columns["beh"] = [rng.integers(0, 32, size=2).tolist() for _ in range(rows)] + os.makedirs(path, exist_ok=True) + pq.write_table( + pa.table( + {k: pa.array(v, type=pa.list_(pa.int64())) for k, v in columns.items()} + ), + os.path.join(path, "part-0.parquet"), + ) + return os.path.join(path, "*.parquet") def export_tiny_genrec( test_dir: str, projected: bool = True, bundle_uuid: Optional[str] = None, - hidden_size: int = 32, + use_dense_ema: bool = False, + env: Optional[Dict[str, str]] = None, ) -> ExportedGenrec: """Train nothing, save a checkpoint of a tiny genrec model, and export it. The prompt carries an INLINE SID history and, when ``projected``, one - PROJECTED behaviour slot, which is what makes the export write a front-end - with tables and a projection. + PROJECTED behaviour slot, which is what makes the export carry a table and + a projection. The export traces on one batch of mock parquet, as any tzrec + export does. Args: test_dir: scratch directory. projected: whether to add the projected slot. bundle_uuid: when given, a SID manifest carrying it is written and referenced, so the export records a bundle identity. - hidden_size: the tiny backbone's hidden size. + use_dense_ema: ask the export for Dense EMA weights, which the HF + export refuses. + env: extra environment for the export, e.g. ``USE_DISTRIBUTED_EMBEDDING``. Returns: Everything a test needs to read the export back. """ + from unittest import mock + + from tzrec.constant import Mode + from tzrec.datasets.dataset import create_dataloader + backbone = os.path.join(test_dir, "backbone") create_tiny_causal_lm(64).save_pretrained(backbone) tok = create_prompt_tokenizer(os.path.join(test_dir, "tok.json"), _WORDS) @@ -242,11 +287,13 @@ def export_tiny_genrec( manifest = os.path.join(test_dir, "manifest.json") with open(manifest, "w") as f: json.dump({"codebook": _CODEBOOK, "bundle_uuid": bundle_uuid}, f) + data_glob = _write_mock_prompt_data(os.path.join(test_dir, "data"), projected) config = EasyRecConfig() text_format.Merge( f''' -train_input_path: "" eval_input_path: "" model_dir: "{test_dir}/train" +train_input_path: "{data_glob}" eval_input_path: "{data_glob}" +model_dir: "{test_dir}/train" train_config {{ sparse_optimizer {{ adagrad_optimizer {{ lr: 0.0 }} constant_learning_rate {{}} }} dense_optimizer {{ adam_optimizer {{ lr: 0.0001 }} constant_learning_rate {{}} }} @@ -256,6 +303,7 @@ def export_tiny_genrec( batch_size: 4 dataset_type: ParquetDataset fg_mode: FG_NONE label_fields: "answer" num_workers: 1 }} +{"export_config { use_dense_ema: true }" if use_dense_ema else ""} feature_configs {{ {_HIST} }} {"feature_configs { " + projected_feature("beh", 8) + " }" if projected else ""} prompt_config {{ @@ -291,7 +339,34 @@ def export_tiny_genrec( model, os.path.join(config.model_dir, "model.ckpt-1") ) export_dir = os.path.join(test_dir, "export") - export(config_path, export_dir, checkpoint_path=checkpoint_dir) + # export_model_normal opens a single-process gloo group; close it again so + # a later test in this process can open its own + group_was_open = dist.is_initialized() + with mock.patch.dict( + os.environ, + { + "MASTER_ADDR": "127.0.0.1", + "MASTER_PORT": os.environ.get("MASTER_PORT", str(_free_port())), + "RANK": "0", + "LOCAL_RANK": "0", + "WORLD_SIZE": "1", + **(env or {}), + }, + ): + try: + export(config_path, export_dir, checkpoint_path=checkpoint_dir) + finally: + if dist.is_initialized() and not group_was_open: + dist.destroy_process_group() + dataloader = create_dataloader( + config.data_config, features, data_glob, mode=Mode.PREDICT + ) + batch = next(iter(dataloader)) + sample_data = { + k: v + for k, v in batch.to_dict(sparse_dtype=torch.int64).items() + if not k.startswith(PROMPT_INFO_PREFIX) + } return ExportedGenrec( config=config, config_path=config_path, @@ -299,4 +374,14 @@ def export_tiny_genrec( compiled_prompt=compiled_prompt, checkpoint_dir=checkpoint_dir, export_dir=export_dir, + sample_data=sample_data, ) + + +def _free_port() -> int: + """A port the single-process gloo group can bind.""" + import socket + + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + return int(sock.getsockname()[1]) diff --git a/tzrec/utils/export_util.py b/tzrec/utils/export_util.py index 7383c2ce..7ffb06f0 100644 --- a/tzrec/utils/export_util.py +++ b/tzrec/utils/export_util.py @@ -72,6 +72,7 @@ from tzrec.utils.dist_util import DistributedModelParallel, init_process_group from tzrec.utils.filesystem_util import url_to_fs from tzrec.utils.fx_util import ( + UNTRACEABLE_MODULES, fx_mark_keyed_tensor, fx_mark_seq_ec_jt, fx_mark_seq_len, @@ -375,6 +376,24 @@ def export_model_normal( if assets is not None: for asset in assets: shutil.copy(asset, save_dir) + _export_extra_assets(model, pipeline_config, checkpoint_path, save_dir) + + +def _export_extra_assets( + model: nn.Module, + pipeline_config: EasyRecConfig, + checkpoint_path: str, + save_dir: str, +) -> None: + """Let a served module write what its runtime reads beside the scripted model. + + Discovered by duck typing through the inference wrappers, the way + checkpointing finds ``hf_backbone``; runs inside the export's own save dir + so a remote export dir is uploaded together with the scripted model. + """ + served = checkpoint_util.unwrap_to(model, "export_assets") + if served is not None: + served.export_assets(pipeline_config, checkpoint_path, save_dir) def _prepare_single_rank_distributed_embedding_export() -> bool: @@ -434,6 +453,15 @@ def _get_leaf_module_names_helper( return list(leaf_module_names) +def _get_untraceable_leaf_module_names(model: torch.nn.Module) -> List[str]: + """Paths of the modules FX cannot trace, for a tracer that matches by path.""" + return [ + path + for path, module in model.named_modules() + if type(module).__name__ in UNTRACEABLE_MODULES + ] + + def _get_dense_embedding_leaf_module_names(model: torch.nn.Module) -> List[str]: """Get dense-embedding modules to keep as FX leaf modules during export. @@ -1666,7 +1694,10 @@ def export_distributed_embedding( torch.cuda.empty_cache() unwrap_model = dmp_model.module - tracer = Tracer(leaf_modules=_get_sharded_leaf_module_names(unwrap_model)) + tracer = Tracer( + leaf_modules=_get_sharded_leaf_module_names(unwrap_model) + + _get_untraceable_leaf_module_names(unwrap_model) + ) full_graph = tracer.trace(unwrap_model) if is_rank_zero: @@ -1871,6 +1902,7 @@ def export_distributed_embedding( merged_emb_json = _merge_sharded_embedding_json(emb_json_files) with open(os.path.join(save_dir_sparse, "sparse_embedding.json"), "w") as f: json.dump(merged_emb_json, f, indent=4) + _export_extra_assets(model, pipeline_config, checkpoint_path, save_dir) class _SparseMarkCapture(Interpreter): diff --git a/tzrec/utils/fx_util.py b/tzrec/utils/fx_util.py index d504234c..01099e54 100644 --- a/tzrec/utils/fx_util.py +++ b/tzrec/utils/fx_util.py @@ -15,6 +15,11 @@ from torchrec import JaggedTensor, KeyedTensor from torchrec.fx import symbolic_trace as _symbolic_trace +# Modules whose forward FX cannot record -- they branch on tensor values or +# turn them into Python ints -- so tracing keeps them opaque and TorchScript +# compiles them whole. Matched by class name. +UNTRACEABLE_MODULES = ["ComputeJTDictToKJT", "PromptAssembler"] + def symbolic_trace( # pyre-ignore[24] @@ -39,8 +44,7 @@ def symbolic_trace( Returns: GraphModule: a Module created from the recorded operations from ``root``. """ - # ComputeJTDictToKJT could not be traced - _leaf_modules = ["ComputeJTDictToKJT"] + _leaf_modules = list(UNTRACEABLE_MODULES) if leaf_modules: _leaf_modules.extend(leaf_modules) return _symbolic_trace(root, concrete_args, _leaf_modules) diff --git a/tzrec/utils/hf_export_util.py b/tzrec/utils/hf_export_util.py index 8d3279cf..44095f57 100644 --- a/tzrec/utils/hf_export_util.py +++ b/tzrec/utils/hf_export_util.py @@ -18,7 +18,7 @@ import json import os import shutil -from typing import Dict, Optional, Set +from typing import Any, Dict, Optional, Set import torch from safetensors.torch import save_file @@ -28,6 +28,9 @@ from tzrec.utils import checkpoint_util from tzrec.utils.logging_util import logger +SERVING_ARCH = "PromptGenRecForCausalLM" +SERVING_MODEL_TYPE = "prompt_genrec" + _HF_ASSET_FILES = ( "config.json", "generation_config.json", @@ -75,6 +78,39 @@ def write_hf_assets(wrapped_model: nn.Module, save_dir: str) -> None: json.dump(meta, f, indent=2) +def write_composite_config(export_dir: str) -> None: + """Rewrite ``config.json`` so the backbone sits under ``text_config``. + + That is what names the backbone to a serving runtime that composes an + arbitrary causal LM behind one registered architecture, and what makes it + treat the model as carrying a second modality, which is what a projected + prompt slot is. + + Args: + export_dir: the HuggingFace export directory. + """ + path = os.path.join(export_dir, "config.json") + with open(path, "r") as f: + backbone: Dict[str, Any] = json.load(f) + if backbone.get("model_type") == SERVING_MODEL_TYPE: + return + composite: Dict[str, Any] = { + "architectures": [SERVING_ARCH], + "model_type": SERVING_MODEL_TYPE, + "text_config": backbone, + } + # a runtime that reads only the outer config still needs to size its cache + for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): + if key in backbone: + composite[key] = backbone[key] + with open(path, "w") as f: + json.dump(composite, f, indent=2) + logger.info( + f"wrote a composite config naming backbone " + f"{backbone.get('architectures', ['?'])[0]} under {SERVING_ARCH}." + ) + + def dcp_to_hf(ckpt_dir: str, out_dir: str) -> None: """Convert a checkpoint with co-located HF assets to a ``from_pretrained`` dir. From 17504f62b2e0d72fb03293111a8bc91fdfc07e36 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 19:04:29 +0800 Subject: [PATCH 04/21] [refactor] keep only the SID space in prompt.json Nothing in the serving engine reads prompt.json; its one reader is the offline constraint-index builder, which needs the token base, the per-level bands and the bundle identity, all of which the resolved SID space carries. The decode schedule, the prefix bound, the length ceilings, the digests and the front-end description were written for readers that do not exist, so they go, together with the front-end's serving-contract helper. The export contract is still moving, so its user-facing documentation waits until it settles. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- docs/source/usage/export.md | 25 -------- tzrec/models/genrec_model.py | 46 +------------- tzrec/prompt/persist.py | 66 +++++---------------- tzrec/tests/genrec_serving_contract_test.py | 11 +--- tzrec/tests/prompt_integration_test.py | 2 - 5 files changed, 20 insertions(+), 130 deletions(-) diff --git a/docs/source/usage/export.md b/docs/source/usage/export.md index 0369ba1c..734cf420 100644 --- a/docs/source/usage/export.md +++ b/docs/source/usage/export.md @@ -221,28 +221,3 @@ eascmd -i ${ACCESS_KEY_ID} -k ${ACCESS_KEY_SECRET} -e ${ENDPOINT} create aot_exp ``` 任务运行结束后,`--export_dir` 指向的目录即为导出好的模型,将其作为在线服务的模型路径部署即可(部署方式参见 [模型服务](serving.md))。 - -(genrec-export)= - -## 生成式推荐模型(genrec)导出 - -`genrec_causal_lm_model` 与其他模型一样通过 `tzrec.export` 导出:tzrec 侧导出的是 **prompt 前端**(特征 -> 拼接后的 token 流与 PROJECTED slot 的投影向量),量化、`INPUT_TILE`、`USE_DISTRIBUTED_EMBEDDING` 等环境变量与普通模型一致;LLM 骨干网络则以 HuggingFace 权重的形式放在同一目录下,交给 SGLang 等 LLM 推理引擎解码: - -``` -export_dir/ - scripted_model.pt # prompt 前端:特征 -> input_ids / hole_positions / slot_embeds / hole_keys / hole_slot_counts - fg.json pipeline.config model_acc.json - dense_meta.json sparse/ # 仅 USE_DISTRIBUTED_EMBEDDING=1 时,与普通模型的分布式 embedding 导出一致 - config.json # 复合结构:architectures 为 PromptGenRecForCausalLM,骨干网络配置位于 text_config - model.safetensors # 骨干网络权重,参数名与骨干网络一致 - generation_config.json - prompt/ - prompt.json # 服务契约:sid_space(band、base_vocab_size、bundle_uuid)、decode 调度、vocab_hash/plan_hash、frontend 输入输出 - tokenizer/ # 扩展了 SID token 的 tokenizer,可由 AutoTokenizer 加载(对应 SGLang 的 --tokenizer-path) -``` - -- TorchEasyRec Processor 以 `export_dir` 作为 `model_path` 加载前端;SGLang 以 `--model-path export_dir` 加载骨干网络,并通过 `SGLANG_PROMPT_FRONTEND_PATH=export_dir/scripted_model.pt` 加载前端。 -- 默认导出下前端内置 PROJECTED slot 的 embedding 表,输入为 FG 输出的原始 id(`{feature}.values` / `.lengths` / `.key_lengths`)。设置 `USE_DISTRIBUTED_EMBEDDING=1` 时,前端改为读取 Processor 分布式 embedding 阶段查表后的向量(命名由 `dense_meta.json` 描述),embedding 表以 `sparse/*.npz` 导出。此时前端仍需要 INLINE slot 与 PROJECTED 成员的原始 id(用于计算 `hole_keys`),`prompt.json` 的 `frontend.inputs` 列出了全部输入 key。 -- 前端仅支持 TorchScript 导出:prompt 拼接的形状随请求变化,`ENABLE_AOT` / `ENABLE_TRT` / `USE_RTP` 不适用。 -- 约束解码索引不由 tzrec 生成:推理侧根据 `prompt/prompt.json` 与 SID bundle 的 `sid_to_items` 构建(SGLang 侧 `python -m sglang.srt.beam_search.build_constraint_csr`),并以 `bundle_uuid` 校验索引与模型是否来自同一 bundle。 -- genrec 导出不支持 `export_config.use_dense_ema=true`。 diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 87a9092c..3cf80e25 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -51,7 +51,7 @@ ) from tzrec.prompt.compile import compile_prompt from tzrec.prompt.persist import PROMPT_DIR, TOKENIZER_DIR, write_serving_contract -from tzrec.prompt.types import CompiledPrompt, FillMode, PromptPlan, SlotSeg +from tzrec.prompt.types import CompiledPrompt, PromptPlan from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig from tzrec.protos.pipeline_pb2 import EasyRecConfig @@ -60,9 +60,6 @@ from tzrec.utils.logging_util import logger SLOT_EMBEDS = "slot_embeds" -SCRIPTED_MODEL_FILENAME = "scripted_model.pt" -LOOKUP_ARTIFACT = "artifact" -LOOKUP_HOST = "host" _PARAM_DTYPE: Dict[int, torch.dtype] = { GenrecModelConfig.FP32: torch.float32, @@ -436,50 +433,13 @@ def predict(self, batch: Batch) -> Dict[str, torch.Tensor]: ) return out - def _serving_contract(self) -> Dict[str, Any]: - """What the exported module reads and writes, for ``prompt.json``.""" - by_name = {feature.name: feature for feature in self._features} - inputs: List[str] = [] - plan = self._prompt.prompt_plan - for seg in plan.segments: - if not isinstance(seg, SlotSeg): - continue - members = ( - seg.feature_names[:1] - if seg.fill is FillMode.INLINE - else seg.feature_names - ) - for name in members: - feature = by_name[name] - inputs.extend([f"{name}.values", f"{name}.lengths"]) - if feature.is_sequence and feature.value_dim != 1: - inputs.append(f"{name}.key_lengths") - return { - "model": SCRIPTED_MODEL_FILENAME, - "lookup": ( - LOOKUP_HOST - if acc_utils.use_distributed_embedding() - else LOOKUP_ARTIFACT - ), - "inputs": list(dict.fromkeys(inputs)), - "outputs": [ - INPUT_IDS, - CU_SEQLENS, - HOLE_POSITIONS, - HOLE_KEYS, - HOLE_SLOT_COUNTS, - SLOT_EMBEDS, - ], - } - def export_assets( self, pipeline_config: EasyRecConfig, checkpoint_path: str, save_dir: str ) -> None: """Write what an LLM engine reads beside the scripted front-end. The HuggingFace weights and composite config, the extended tokenizer and - the serving contract. Called by the export on rank 0, inside its save - dir. + the SID space. Called by the export on rank 0, inside its save dir. Args: pipeline_config: the pipeline being exported. @@ -502,4 +462,4 @@ def export_assets( list(pipeline_config.data_config.label_fields), tokenizer_dir=os.path.join(save_dir, PROMPT_DIR, TOKENIZER_DIR), ) - write_serving_contract(self._prompt, save_dir, self._serving_contract()) + write_serving_contract(self._prompt, save_dir) diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py index 81b003af..f2009dfb 100644 --- a/tzrec/prompt/persist.py +++ b/tzrec/prompt/persist.py @@ -9,23 +9,25 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Checks a checkpoint's prompt contract on restore, and publishes it on export. +"""Checks a checkpoint's prompt contract on restore, and publishes the SID space. The digests ride in the HF export metadata the checkpoint already carries, so there is no second file that can drift from the weights beside it. Export -additionally writes ``prompt/prompt.json``: what a serving runtime reads, and -only that. The plan itself is deliberately not published there -- it reaches -serving compiled into the front-end artifact, and a copy a runtime could -interpret would invite the second assembler this design exists to prevent. +additionally writes ``prompt/prompt.json`` with the resolved SID space: the +token base, the per-level bands and the bundle identity a serving side needs +to build its constraint index and to refuse one built from another bundle. +The plan itself is deliberately not published -- it reaches serving compiled +into the front-end, and a copy a runtime could interpret would invite the +second assembler this design exists to prevent. """ import dataclasses import json import os -from typing import Any, Dict, List, Optional +from typing import Dict, Optional from tzrec.constant import HF_EXPORT_META_FILENAME -from tzrec.prompt.types import CompiledPrompt, SlotSeg, Static +from tzrec.prompt.types import CompiledPrompt from tzrec.utils.logging_util import logger PROMPT_DIR = "prompt" @@ -33,62 +35,26 @@ TOKENIZER_DIR = "tokenizer" -def _decode_schedule(compiled_prompt: CompiledPrompt) -> List[Dict[str, int]]: - """One entry per response position: a forced token or the band it is masked to. - - Derived from the response segments, so a runtime gets the schedule without - gaining the ability to interpret a plan. - """ - sid_space = compiled_prompt.sid_space - schedule: List[Dict[str, int]] = [] - for seg in compiled_prompt.prompt_plan.response_segments: - if isinstance(seg, Static): - schedule.extend({"token_id": int(t)} for t in seg.token_ids) - continue - assert isinstance(seg, SlotSeg) and sid_space is not None - width = seg.width.num_positions or 0 - for index in range(width): - level = index % sid_space.num_levels - schedule.append( - { - "level": level, - "band_lo": sid_space.band_lo[level], - "band_hi": sid_space.band_hi[level], - } - ) - return schedule - - -def write_serving_contract( - compiled_prompt: CompiledPrompt, export_dir: str, frontend: Dict[str, Any] -) -> str: - """Write ``prompt/prompt.json``, everything a serving runtime reads. +def write_serving_contract(compiled_prompt: CompiledPrompt, export_dir: str) -> str: + """Write ``prompt/prompt.json``: the resolved SID space. Args: compiled_prompt: the compiled prompt. - export_dir: the HuggingFace export directory. - frontend: how the exported front-end is laid out and what it reads. + export_dir: the export directory. Returns: The path written. """ - plan = compiled_prompt.prompt_plan out = os.path.join(export_dir, PROMPT_DIR) os.makedirs(out, exist_ok=True) path = os.path.join(out, PROMPT_CONTRACT_FILENAME) + sid_space = compiled_prompt.sid_space with open(path, "w") as f: json.dump( { - "sid_space": dataclasses.asdict(compiled_prompt.sid_space) - if compiled_prompt.sid_space is not None - else None, - "decode": _decode_schedule(compiled_prompt), - "static_prefix_len": plan.static_prefix_len, - "max_length": plan.max_length, - "max_total_length": plan.max_total_length, - "vocab_hash": compiled_prompt.vocab_hash, - "plan_hash": compiled_prompt.plan_hash, - "frontend": frontend, + "sid_space": dataclasses.asdict(sid_space) + if sid_space is not None + else None }, f, indent=2, diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index 2a877b92..e9569a7c 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -138,6 +138,7 @@ def test_host_item_hash_refolds_hole_keys(self) -> None: self.assertNotEqual(_item_hash(keys), _item_hash(keys.flip(0))) def test_prompt_json_carries_the_index_builder_inputs(self) -> None: + self.assertEqual(list(self.contract), ["sid_space"]) space = self.contract["sid_space"] compiled = self.exported.compiled_prompt.sid_space self.assertEqual(space["base_vocab_size"], compiled.base_vocab_size) @@ -145,16 +146,6 @@ def test_prompt_json_carries_the_index_builder_inputs(self) -> None: self.assertEqual(space["band_lo"], list(compiled.band_lo)) self.assertEqual(space["band_hi"], list(compiled.band_hi)) self.assertEqual(space["bundle_uuid"], "bundle-test") - self.assertEqual(len(self.contract["decode"]), 3) - self.assertEqual( - self.contract["vocab_hash"], self.exported.compiled_prompt.vocab_hash - ) - self.assertEqual(self.contract["frontend"]["model"], "scripted_model.pt") - self.assertEqual(self.contract["frontend"]["lookup"], "artifact") - self.assertIn("beh.values", self.contract["frontend"]["inputs"]) - self.assertEqual( - self.contract["plan_hash"], self.exported.compiled_prompt.plan_hash - ) def test_tokenizer_dir_decodes_a_sid_atom(self) -> None: from transformers import AutoTokenizer diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index 430dc13b..32b1bfc9 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -183,8 +183,6 @@ def test_distributed_embedding_export_writes_the_processor_shape(self) -> None: with open(os.path.join(dist_dir, "dense_meta.json"), "r") as f: dense_meta = json.load(f) self.assertEqual(dense_meta["sequence__ec"], ["beh__ec", "beh__lengths"]) - with open(os.path.join(dist_dir, "prompt", "prompt.json"), "r") as f: - self.assertEqual(json.load(f)["frontend"]["lookup"], "host") # a processor simulator: one request is one user, looked up in the # exported tables the way the distributed-embedding stage does, then From 6f5869f0aed32690d162c02aed6c26b0a7d4365b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 20:04:31 +0800 Subject: [PATCH 05/21] [refactor] move the hole_keys fold out of the prompt assembler The assembler is prompt structure only. The fold is a serving-only module composed by ScriptWrapper, so training no longer computes keys it discards. PromptPlan loses its fold constants and the assembler its host-ABI fallback. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/datasets/dataset.py | 1 - tzrec/models/genrec_model.py | 13 +- tzrec/models/genrec_model_test.py | 21 +- tzrec/models/model.py | 14 +- tzrec/prompt/assembler.py | 227 ++-------------- tzrec/prompt/assembler_test.py | 150 +---------- tzrec/prompt/hole_keys.py | 220 ++++++++++++++++ tzrec/prompt/hole_keys_test.py | 275 ++++++++++++++++++++ tzrec/prompt/persist.py | 9 +- tzrec/prompt/types.py | 36 +-- tzrec/tests/genrec_serving_contract_test.py | 6 +- tzrec/tests/prompt_integration_test.py | 16 +- tzrec/tests/prompt_test_util.py | 8 +- tzrec/utils/fx_util.py | 2 +- 14 files changed, 564 insertions(+), 434 deletions(-) create mode 100644 tzrec/prompt/hole_keys.py create mode 100644 tzrec/prompt/hole_keys_test.py diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index bc851322..c12f406d 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -117,7 +117,6 @@ def __init__( PromptAssembler( compiled_prompt.prompt_plan, compiled_prompt.sid_space, - plan_hash=compiled_prompt.plan_hash, include_response=mode != Mode.PREDICT, ) if compiled_prompt is not None diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 3cf80e25..0231bcbd 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -39,17 +39,16 @@ from tzrec.modules.prompt_projection import PromptProjection from tzrec.prompt.assembler import ( CU_SEQLENS, - HOLE_KEYS, HOLE_POSITIONS, HOLE_SLOT_COUNTS, INPUT_IDS, PROMPT_CU_SEQLENS, - PROMPT_HOLE_KEYS, PROMPT_HOLE_POSITIONS, PROMPT_HOLE_SLOT_COUNTS, PROMPT_INPUT_IDS, ) from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.hole_keys import HOLE_KEYS, PROMPT_HOLE_KEYS from tzrec.prompt.persist import PROMPT_DIR, TOKENIZER_DIR, write_serving_contract from tzrec.prompt.types import CompiledPrompt, PromptPlan from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig @@ -101,11 +100,6 @@ def __init__( f"compile_prompt(pipeline_config.prompt_config, features) and " f"pass it to _create_model." ) - if compiled_prompt.sid_space is None: - raise ValueError( - f"{type(self).__name__}: prompt_config declares no sid_space, " - f"so there is no SID vocabulary to extend or decode." - ) if compiled_prompt.prompt_plan.logits_suffix_len is None: raise ValueError( f"{type(self).__name__}: the response is unbounded, so the " @@ -368,6 +362,11 @@ class GenrecFrontEnd(nn.Module): projected slot embeddings; an LLM engine gathers the LM's own table, scatters ``slot_embeds`` at ``hole_positions`` and decodes. + The walk and the ``hole_keys`` fold both read the parsed feature dict, so + under distributed embedding the dense stage must still receive every slot + member's raw ``.values`` / ``.lengths`` / ``.key_lengths`` beside the + looked-up embeddings. + Args: model: the genrec model to serve. """ diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 0b1b170c..402490f6 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -9,7 +9,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import dataclasses import json import os import shutil @@ -29,7 +28,6 @@ ) from tzrec.models.model import ScriptWrapper from tzrec.prompt.assembler import ( - HOLE_KEYS, HOLE_POSITIONS, INPUT_IDS, PROMPT_HOLE_POSITIONS, @@ -37,6 +35,7 @@ PromptAssembler, ) from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig @@ -80,14 +79,6 @@ def test_rejects_a_model_built_without_a_prompt(self) -> None: with self.assertRaisesRegex(ValueError, "needs a compiled prompt"): self._model(compiled_prompt=None) - def test_rejects_a_prompt_that_declares_no_sid_space(self) -> None: - # compile_prompt refuses this config, so the model precondition is - # reachable only by constructing a prompt directly - compiled_prompt = dataclasses.replace(self.compiled_prompt, sid_space=None) - - with self.assertRaisesRegex(ValueError, "declares no sid_space"): - self._model(compiled_prompt=compiled_prompt) - def test_shared_projection_name_requires_matching_widths(self) -> None: features = [ create_prompt_feature(_HIST), @@ -263,13 +254,13 @@ def test_front_end_returns_the_walk_and_the_projected_slots(self) -> None: out = self.wrapped(self.data) compiled = self.compiled_prompt walk = PromptAssembler( - compiled.prompt_plan, - compiled.sid_space, - plan_hash=compiled.plan_hash, - include_response=False, + compiled.prompt_plan, compiled.sid_space, include_response=False )(self.data) - for key in (INPUT_IDS, HOLE_POSITIONS, HOLE_KEYS): + for key in (INPUT_IDS, HOLE_POSITIONS): self.assertTrue(torch.equal(out[key], walk[key]), key) + self.assertTrue( + torch.equal(out[HOLE_KEYS], HoleKeyBuilder(compiled.prompt_plan)(self.data)) + ) batch = self.wrapped.get_batch(self.data) expected = project_slots( self.model.embedding_group, diff --git a/tzrec/models/model.py b/tzrec/models/model.py index b9eed9ef..06bdb6a8 100644 --- a/tzrec/models/model.py +++ b/tzrec/models/model.py @@ -31,6 +31,7 @@ from tzrec.loss.pe_mtl_loss import ParetoEfficientMultiTaskLoss from tzrec.modules.utils import BaseModule from tzrec.prompt.assembler import OUTPUT_KEYS, PROMPT_INFO_PREFIX, PromptAssembler +from tzrec.prompt.hole_keys import PROMPT_HOLE_KEYS, HoleKeyBuilder from tzrec.protos.loss_pb2 import LossConfig from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.utils import config_util @@ -392,7 +393,8 @@ class ScriptWrapper(BaseModule): A module that exposes ``compiled_prompt`` gets its prompt assembled here, from the parsed dict, exactly as the training collator does: one walk, two - call sites. + call sites. Serving alone also folds the prefix-cache ``hole_keys`` here, + which training never computes. """ def __init__(self, module: nn.Module) -> None: @@ -407,14 +409,14 @@ def __init__(self, module: nn.Module) -> None: prompt = getattr(module, "compiled_prompt", None) self._prompt_assembler = ( PromptAssembler( - prompt.prompt_plan, - prompt.sid_space, - plan_hash=prompt.plan_hash, - include_response=False, + prompt.prompt_plan, prompt.sid_space, include_response=False ) if prompt is not None else None ) + self._hole_keys = ( + HoleKeyBuilder(prompt.prompt_plan) if prompt is not None else None + ) @property def features(self) -> List[BaseFeature]: @@ -439,6 +441,8 @@ def get_batch( batch.additional_infos.update( {PROMPT_INFO_PREFIX + k: streams[k] for k in OUTPUT_KEYS} ) + if self._hole_keys is not None: + batch.additional_infos[PROMPT_HOLE_KEYS] = self._hole_keys(data) batch = batch.to(device, non_blocking=True) return batch diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index 422952d0..0fb8c00f 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -19,23 +19,17 @@ which is what lets ``torch.jit.script`` carry the same module into a runtime that has no tzrec source. Serving therefore never reimplements this walk. -``hole_keys`` is computed at both call sites and discarded by training. It is -integer end to end: ``int64`` addition is associative and commutative and wraps -deterministically, so the fold cannot depend on the order ``index_add_`` -happens to reduce in, on the device, or on how the batch was split. A float -accumulator would satisfy "do not fold the projected vector" in letter and -reintroduce the variance in spirit. +The walk is prompt structure only. The prefix-cache identity of each hole, +``hole_keys``, is a serving concern that ``hole_keys.py`` folds beside it. """ -import hashlib -from typing import Dict, Final, List, Optional, Tuple +from typing import Dict, Final, List, Tuple import torch from torch import nn from tzrec.prompt.types import ( FillMode, - FoldConstants, PromptPlan, ResolvedSidSpace, SlotSeg, @@ -46,7 +40,6 @@ INPUT_IDS = "input_ids" CU_SEQLENS = "cu_seqlens" HOLE_POSITIONS = "hole_positions" -HOLE_KEYS = "hole_keys" HOLE_SLOT_COUNTS = "hole_slot_counts" RESPONSE_LENGTHS = "response_lengths" MAX_SEQLEN = "max_seqlen" @@ -56,7 +49,6 @@ INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, - HOLE_KEYS, HOLE_SLOT_COUNTS, RESPONSE_LENGTHS, MAX_SEQLEN, @@ -67,26 +59,11 @@ PROMPT_INPUT_IDS = PROMPT_INFO_PREFIX + INPUT_IDS PROMPT_CU_SEQLENS = PROMPT_INFO_PREFIX + CU_SEQLENS PROMPT_HOLE_POSITIONS = PROMPT_INFO_PREFIX + HOLE_POSITIONS -PROMPT_HOLE_KEYS = PROMPT_INFO_PREFIX + HOLE_KEYS PROMPT_HOLE_SLOT_COUNTS = PROMPT_INFO_PREFIX + HOLE_SLOT_COUNTS PROMPT_MAX_SEQLEN = PROMPT_INFO_PREFIX + MAX_SEQLEN PROMPT_RESPONSE_LENGTHS = PROMPT_INFO_PREFIX + RESPONSE_LENGTHS -@torch.jit.script -def mix64(z: torch.Tensor) -> torch.Tensor: - """A SplitMix64-shaped avalanche over int64, wrapping. - - torch's right shift on a signed integer is arithmetic, so every shift the - mixer wants as logical is masked back. Getting that wrong is not a weaker - hash, it is a different function on negative inputs. - """ - z = z * (-7046029254386353131) - z = (z ^ ((z >> 30) & 0x3FFFFFFFF)) * (-4658895280553007687) - z = (z ^ ((z >> 27) & 0x1FFFFFFFFF)) * (-7723592293110705685) - return z ^ ((z >> 31) & 0x1FFFFFFFF) - - @torch.jit.script def _exclusive_cumsum(values: torch.Tensor) -> torch.Tensor: """Exclusive prefix sum along dim 0.""" @@ -126,30 +103,6 @@ def _destinations(seg_start: torch.Tensor, seg_len: torch.Tensor) -> torch.Tenso return torch.repeat_interleave(seg_start, seg_len) + _within_row_index(seg_len) -def _wrap64(value: int) -> int: - """Reduce a Python int into the signed 64-bit range, wrapping. - - Every constant the fold mixes is built on the host and handed to torch as a - scalar, and torch refuses a scalar outside the tensor's dtype. Wrapping here - is what makes "the fold wraps" true at the boundary as well as inside it. - """ - value &= (1 << 64) - 1 - return value - (1 << 64) if value >= 1 << 63 else value - - -def _plan_salt(fold: FoldConstants, plan_hash: str) -> int: - """A 64-bit digest of ``plan_hash``, as a signed multiple of the plan constant.""" - if not plan_hash: - return 0 - digest = hashlib.sha256(plan_hash.encode("utf-8")).digest()[:8] - return _wrap64(int.from_bytes(digest, "little") * fold.plan) - - -def host_lengths_key(feature_name: str) -> str: - """The per-row item count a host lookup stage emits beside its rows.""" - return feature_name + "__lengths" - - class PromptAssembler(nn.Module): """Walks a compiled plan to build one batch's packed token stream. @@ -160,10 +113,7 @@ class PromptAssembler(nn.Module): Args: prompt_plan: the compiled walk order and its constants. - sid_space: resolved SID token space; required when a slot renders SID - codes or is projected. - plan_hash: the compiled plan's hash; its low bits salt every key, so - two plans cannot cross-match in a shared prefix cache. + sid_space: the resolved SID token space. include_response: whether to read and emit the supervised tail. """ @@ -172,9 +122,6 @@ class PromptAssembler(nn.Module): KIND_STATIC: Final[int] = 0 KIND_INLINE: Final[int] = 1 KIND_PROJECTED: Final[int] = 2 - # member index and position within a hole packed into one integer, wide - # enough that a position cannot carry into the member index - MEMBER_STRIDE: Final[int] = 1 << 32 kinds: List[int] names: List[str] @@ -182,7 +129,6 @@ class PromptAssembler(nn.Module): exact_widths: List[int] hole_slots: List[int] is_sequences: List[bool] - salts: List[int] member_names: List[List[str]] level_lo: List[int] level_hi: List[int] @@ -190,8 +136,7 @@ class PromptAssembler(nn.Module): def __init__( self, prompt_plan: PromptPlan, - sid_space: Optional[ResolvedSidSpace] = None, - plan_hash: str = "", + sid_space: ResolvedSidSpace, include_response: bool = True, ) -> None: super().__init__() @@ -200,30 +145,17 @@ def __init__( if include_response: segments = segments + tuple(prompt_plan.response_segments) - inline = [ - s.name - for s in segments - if isinstance(s, SlotSeg) and s.fill is FillMode.INLINE - ] - if inline and sid_space is None: - raise ValueError( - f"prompt slots {inline} render INLINE, which means SID codes, but " - f"no sid_space was compiled." - ) + # compile reserves a sentinel whenever a slot is PROJECTED, so -1 is + # never written self.sentinel = -1 - self.id_shift = 0 - self.num_levels = 1 - self.level_lo = [] - self.level_hi = [] - if sid_space is not None: - if sid_space.sentinel_token_id is not None: - self.sentinel = int(sid_space.sentinel_token_id) - self.id_shift = int(sid_space.base_vocab_size) - self.num_levels = int(sid_space.num_levels) - self.level_lo = [int(o) for o in sid_space.level_offsets] - self.level_hi = [ - int(o + c) for o, c in zip(sid_space.level_offsets, sid_space.codebook) - ] + if sid_space.sentinel_token_id is not None: + self.sentinel = int(sid_space.sentinel_token_id) + self.id_shift = int(sid_space.base_vocab_size) + self.num_levels = int(sid_space.num_levels) + self.level_lo = [int(o) for o in sid_space.level_offsets] + self.level_hi = [ + int(o + c) for o, c in zip(sid_space.level_offsets, sid_space.codebook) + ] self.max_length = int(prompt_plan.max_length) self.kinds = [] @@ -233,9 +165,9 @@ def __init__( self.hole_slots = [] self.is_sequences = [] self.member_names = [] - slot_ids: List[int] = [] # holes are grouped by projected occurrence in emission order, which is - # the order of ``projected_slots`` and of the front-end's projections + # the order of ``projected_slots``, of the front-end's projections and + # of ``hole_keys`` occurrences = 0 for index, seg in enumerate(segments): if isinstance(seg, Static): @@ -248,7 +180,6 @@ def __init__( False, [], ) - slot_ids.append(-1) continue assert isinstance(seg, SlotSeg) is_sequence = seg.group_type == FeatureGroupType.JAGGED_SEQUENCE @@ -267,11 +198,6 @@ def __init__( [seg.feature_names[0]], ) else: - if self.sentinel < 0: - raise ValueError( - f"prompt slot [{seg.name}] is PROJECTED but no sentinel token " - f"was compiled; a hole would be indistinguishable from content." - ) self._append( self.KIND_PROJECTED, seg.name, @@ -282,15 +208,9 @@ def __init__( list(seg.feature_names), ) occurrences += 1 - slot_ids.append(int(seg.slot_id)) self.num_segments = len(self.kinds) self.num_hole_slots = occurrences - fold = prompt_plan.fold - self.fold_value = int(fold.value) - self.fold_index = int(fold.index) - plan_salt = _plan_salt(fold, plan_hash) - self.salts = [_wrap64(fold.slot * slot_id + plan_salt) for slot_id in slot_ids] def _append( self, @@ -314,21 +234,19 @@ def _append( def _lengths( self, batch: Dict[str, torch.Tensor], slot: str, member: str, is_sequence: bool ) -> torch.Tensor: - """Per-row item count, from the parsed dict or from a host lookup stage.""" + """Per-row item count of one member, as the data parser emits it.""" key = member + ".lengths" if key in batch: return batch[key].to(torch.int64) - host_key = host_lengths_key(member) - if host_key in batch: - return batch[host_key].to(torch.int64) if is_sequence: raise ValueError( "prompt slot [" + slot - + "] reads [" - + member - + "], which the parser emitted as a scalar: a slot renders a " - + "sequence of SID codes, so the column must be list." + + "] renders a sequence but the batch has no [" + + key + + "]: the column must be list, and under distributed " + + "embedding the processor must pass the raw parsed features " + + "through beside the looked-up embeddings." ) # a dense member has one row per sample and no lengths return torch.ones( @@ -488,87 +406,6 @@ def _segment( (total,), self.sentinel, dtype=torch.int64, device=seg_len.device ) - def _fold_segment( - self, - batch: Dict[str, torch.Tensor], - index: int, - batch_size: int, - hole_base: int, - keys: torch.Tensor, - ) -> None: - """Mix one projected segment's input values into ``keys``. - - Every member value that produces a hole contributes, discriminated by - slot, by member and by its index inside the hole. Without the last two - a two-member slot with values ``(a, b)`` would match one with - ``(b, a)``, and a permuted multi-value item would match itself - reordered -- both plausible, both wrong, and both silent. - """ - salt = self.salts[index] - name = self.names[index] - members = self.member_names[index] - is_sequence = self.is_sequences[index] - for member_index in range(len(members)): - member = members[member_index] - raw = batch[member + ".values"] - lengths = self._lengths(batch, name, member, is_sequence) - key_length_key = member + ".key_lengths" - - if not is_sequence: - # one hole per row; a dense member contributes its float32 bit - # pattern verbatim, which is the parsed input and not a - # computed reduction, so it is stable for a given request - if raw.is_floating_point(): - width = raw.size(1) - values = ( - raw.to(torch.float32) - .contiguous() - .view(torch.int32) - .to(torch.int64) - .reshape(-1) - & 0xFFFFFFFF - ) - hole = torch.repeat_interleave( - torch.arange(batch_size, dtype=torch.int64, device=raw.device), - torch.full( - (batch_size,), width, dtype=torch.int64, device=raw.device - ), - ) - local = ( - torch.arange(width, dtype=torch.int64, device=raw.device) - .unsqueeze(0) - .expand(batch_size, width) - .reshape(-1) - ) - else: - values = raw.to(torch.int64).reshape(-1) - hole = _row_ids(lengths) - local = _within_row_index(lengths) - else: - if raw.is_floating_point(): - raise ValueError( - "prompt slot [" - + name - + "] member [" - + member - + "] is a dense sequence; the fold has no per-item boundary " - + "for it." - ) - values = raw.to(torch.int64).reshape(-1) - if key_length_key in batch: - key_lengths = batch[key_length_key].to(torch.int64).reshape(-1) - hole = _row_ids(key_lengths) - local = _within_row_index(key_lengths) - else: - hole = torch.arange( - values.numel(), dtype=torch.int64, device=values.device - ) - local = torch.zeros_like(hole) - - local = local + member_index * self.MEMBER_STRIDE - mixed = mix64(values * self.fold_value + local * self.fold_index + salt) - keys.index_add_(0, hole + hole_base, mixed) - def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: """Assemble one batch. @@ -577,11 +414,11 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: ``.lengths`` / ``.key_lengths`` as the data parser emits it. Returns: - ``input_ids``, ``cu_seqlens``, ``hole_positions``, ``hole_keys``, - ``hole_slot_counts``, ``response_lengths`` and ``max_seqlen``. The - two hole-indexed streams are row-aligned: entry ``k`` of each - describes the same hole, grouped by projected occurrence in - emission order, then by sample. + ``input_ids``, ``cu_seqlens``, ``hole_positions``, + ``hole_slot_counts``, ``response_lengths`` and ``max_seqlen``. Holes + are grouped by projected occurrence in emission order, then by + sample; the front-end's ``slot_embeds`` and ``hole_keys`` follow the + same order. """ batch_size = self._batch_size(batch) device = batch_device(batch) @@ -639,13 +476,6 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: slot_counts = torch.zeros(self.num_hole_slots, dtype=torch.int64, device=device) for slot in range(self.num_hole_slots): slot_counts[slot] = hole_parts[slot + 1].numel() - keys = torch.zeros(hole_positions.numel(), dtype=torch.int64, device=device) - hole_base = 0 - for slot in range(self.num_hole_slots): - for i in range(self.num_segments): - if self.hole_slots[i] == slot: - self._fold_segment(batch, i, batch_size, hole_base, keys) - hole_base = hole_base + hole_parts[slot + 1].numel() cu_seqlens = torch.cat( [ @@ -663,7 +493,6 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: "input_ids": out, "cu_seqlens": cu_seqlens.to(torch.int32), "hole_positions": hole_positions, - "hole_keys": keys, "hole_slot_counts": slot_counts, "response_lengths": response_lengths, "max_seqlen": max_seqlen, diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index b121e9d6..f2f6405b 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -18,14 +18,12 @@ from tzrec.prompt import assembler from tzrec.prompt.assembler import ( CU_SEQLENS, - HOLE_KEYS, HOLE_POSITIONS, HOLE_SLOT_COUNTS, INPUT_IDS, MAX_SEQLEN, RESPONSE_LENGTHS, PromptAssembler, - mix64, ) from tzrec.prompt.types import ( FillMode, @@ -134,10 +132,6 @@ def _parsed(inline=None, projected=None) -> dict: return out -def _tensors(raw): - return {key: torch.as_tensor(np.asarray(value)) for key, value in raw.items()} - - class PromptAssemblerTest(unittest.TestCase): def test_inline_sid_gets_the_base_vocab_shift(self) -> None: asm = _asm((Static((7, 8)), _slot("hist", FillMode.INLINE))) @@ -292,11 +286,6 @@ def test_over_long_row_is_an_error_not_a_truncation(self) -> None: with self.assertRaisesRegex(ValueError, "never truncated"): asm(_parsed({"hist": [np.array([1, 6, 11])]})) - def test_inline_without_a_sid_space_is_rejected_at_construction(self) -> None: - plan = _plan((_slot("hist", FillMode.INLINE),)) - with self.assertRaisesRegex(ValueError, "no sid_space was compiled"): - PromptAssembler(plan, None) - def test_column_shaped_values_are_flattened(self) -> None: # the data parser emits (total, value_dim) for a dense sequence feature from tzrec.prompt.types import CompiledPrompt, ProjectionPlan @@ -386,7 +375,7 @@ def test_deep_projected_members_emit_one_hole_per_sample(self) -> None: def test_output_keys_match_the_module_constants(self) -> None: """``forward`` writes literals; they must equal the exported names.""" - module = PromptAssembler(_plan((Static((7,)),))) + module = PromptAssembler(_plan((Static((7,)),)), _sid_space()) out = module({"batch_size": torch.tensor(2)}) self.assertEqual( set(out.keys()), @@ -394,7 +383,6 @@ def test_output_keys_match_the_module_constants(self) -> None: INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, - HOLE_KEYS, HOLE_SLOT_COUNTS, RESPONSE_LENGTHS, MAX_SEQLEN, @@ -413,7 +401,7 @@ def test_scripting_preserves_every_output(self) -> None: ), response=(_slot("answer", FillMode.INLINE),), ) - module = PromptAssembler(plan, _sid_space(), plan_hash="a1b2c3d4e5f60718") + module = PromptAssembler(plan, _sid_space()) batch = { "hist.values": torch.tensor([1, 6, 11, 2, 7, 10]), "hist.lengths": torch.tensor([1, 1]), @@ -437,139 +425,5 @@ def test_a_scripted_walk_reports_its_validation(self) -> None: scripted(_parsed({"hist": [np.array([1, 6, 11])]})) -class MixTest(unittest.TestCase): - def test_mix64_matches_a_host_reference(self) -> None: - """The mixer's masked shifts must reproduce SplitMix64 on negatives.""" - - def reference(value: int) -> int: - mask = (1 << 64) - 1 - z = (value * 0x9E3779B97F4A7C15) & mask - z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & mask - z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & mask - z = z ^ (z >> 31) - return z - (1 << 64) if z >= 1 << 63 else z - - values = [0, 1, -1, 2**31, -(2**31), 123456789, -987654321] - got = mix64(torch.tensor(values, dtype=torch.int64)).tolist() - self.assertEqual(got, [reference(v) for v in values]) - - -class FoldTest(unittest.TestCase): - def setUp(self) -> None: - self.beh = _slot("beh", FillMode.PROJECTED, 30) - self.plan = _plan((self.beh,)) - self.module = PromptAssembler( - self.plan, _sid_space(), plan_hash="a1b2c3d4e5f60718" - ) - - def _keys(self, values, lengths): - return self.module( - _tensors( - { - "beh.values": np.array(values, dtype=np.int64), - "beh.lengths": np.array(lengths, dtype=np.int64), - } - ) - )[HOLE_KEYS] - - def test_equal_content_folds_equal(self) -> None: - """A key is a function of the hole's inputs and nothing else.""" - self.assertTrue( - torch.equal(self._keys([5, 6, 7], [3]), self._keys([5, 6, 7], [3])) - ) - - def test_different_content_folds_apart(self) -> None: - """The whole point: a changed input must not reuse the cached KV.""" - first = self._keys([5, 6, 7], [3]) - second = self._keys([5, 6, 8], [3]) - self.assertEqual(first[:2].tolist(), second[:2].tolist()) - self.assertNotEqual(int(first[2]), int(second[2])) - - def test_the_plan_hash_salts_the_keys(self) -> None: - """Two plans must not cross-match in a shared prefix cache.""" - other = PromptAssembler(self.plan, _sid_space(), plan_hash="ffffffffffffffff") - batch = _tensors( - { - "beh.values": np.array([5, 6, 7], dtype=np.int64), - "beh.lengths": np.array([3], dtype=np.int64), - } - ) - self.assertFalse( - torch.equal(self.module(batch)[HOLE_KEYS], other(batch)[HOLE_KEYS]) - ) - - def test_a_permuted_multi_value_item_does_not_collide(self) -> None: - """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" - - def keys(values): - return self.module( - _tensors( - { - "beh.values": np.array(values, dtype=np.int64), - "beh.lengths": np.array([1], dtype=np.int64), - "beh.key_lengths": np.array([3], dtype=np.int64), - } - ) - )[HOLE_KEYS] - - self.assertNotEqual(keys([1, 5, 9]).tolist(), keys([9, 5, 1]).tolist()) - - def test_two_members_exchanging_values_do_not_collide(self) -> None: - """Without the member index a two-member slot is order-blind.""" - slot = _slot("pair", FillMode.PROJECTED, 30, feature_names=("a", "b")) - module = PromptAssembler( - _plan((slot,)), _sid_space(), plan_hash="a1b2c3d4e5f60718" - ) - - def keys(first, second): - return module( - _tensors( - { - "a.values": np.array(first, dtype=np.int64), - "a.lengths": np.array([1], dtype=np.int64), - "b.values": np.array(second, dtype=np.int64), - "b.lengths": np.array([1], dtype=np.int64), - } - ) - )[HOLE_KEYS] - - self.assertNotEqual(keys([3], [9]).tolist(), keys([9], [3]).tolist()) - - def test_a_dense_member_folds_its_bit_pattern(self) -> None: - """A float member contributes the parsed input verbatim, so it is stable.""" - slot = _slot("vec", FillMode.PROJECTED, group_type=FeatureGroupType.DEEP) - module = PromptAssembler(_plan((slot,)), _sid_space(), plan_hash="a1b2") - - def keys(rows): - return module({"vec.values": torch.tensor(rows, dtype=torch.float32)})[ - HOLE_KEYS - ] - - self.assertEqual(keys([[0.5, 1.0]]).tolist(), keys([[0.5, 1.0]]).tolist()) - self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) - - def test_a_dense_sequence_member_is_rejected(self) -> None: - with self.assertRaisesRegex(ValueError, "no per-item boundary"): - self.module( - { - "beh.values": torch.tensor([[0.5], [1.0]]), - "beh.lengths": torch.tensor([2]), - } - ) - - @unittest.skipIf(not torch.cuda.is_available(), "no GPU") - def test_the_fold_is_bit_identical_across_devices(self) -> None: - """Integer addition cannot depend on the order a device reduces in.""" - batch = _tensors( - { - "beh.values": np.arange(64, dtype=np.int64), - "beh.lengths": np.array([32, 32], dtype=np.int64), - } - ) - on_cpu = self.module(batch)[HOLE_KEYS] - on_gpu = self.module({k: v.cuda() for k, v in batch.items()})[HOLE_KEYS] - self.assertTrue(torch.equal(on_cpu, on_gpu.cpu())) - - if __name__ == "__main__": unittest.main() diff --git a/tzrec/prompt/hole_keys.py b/tzrec/prompt/hole_keys.py new file mode 100644 index 00000000..4f3a3f13 --- /dev/null +++ b/tzrec/prompt/hole_keys.py @@ -0,0 +1,220 @@ +# 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. + +"""The ``hole_keys`` fold: a serving-only prefix-cache identity per hole. + +An LLM engine keys its prefix cache on token ids, and a PROJECTED slot writes +the same sentinel at the same position for every request, so the engine needs +a per-hole identity to match on instead. ``HoleKeyBuilder`` folds each hole's +*input* values -- never the projected vector -- into one ``int64`` key per +hole, row-aligned with the assembler's ``hole_positions``: the projected +occurrences in emission order, then the samples. Only the serving wrapper +composes it; training never computes a key it would discard. + +The fold is integer end to end: ``int64`` addition is associative and +commutative and wraps deterministically, so a key cannot depend on the order +``index_add_`` happens to reduce in, on the device, or on how the batch was +split. A float accumulator would satisfy "do not fold the projected vector" in +letter and reintroduce the variance in spirit. + +No plan or model identity enters a key. A cache shared across model versions +already matches their static tokens alike, so only the engine can namespace +it, and the engine folds the keys it receives with a position constant of its +own. Nothing outside the scripted front-end reads the constants below. +""" + +from typing import Dict, Final, List + +import torch +from torch import nn + +from tzrec.prompt.assembler import ( + PROMPT_INFO_PREFIX, + _row_ids, + _within_row_index, + batch_device, +) +from tzrec.prompt.types import PromptPlan +from tzrec.protos.model_pb2 import FeatureGroupType + +HOLE_KEYS = "hole_keys" +PROMPT_HOLE_KEYS = PROMPT_INFO_PREFIX + HOLE_KEYS + + +@torch.jit.script +def mix64(z: torch.Tensor) -> torch.Tensor: + """A SplitMix64-shaped avalanche over int64, wrapping. + + torch's right shift on a signed integer is arithmetic, so every shift the + mixer wants as logical is masked back. Getting that wrong is not a weaker + hash, it is a different function on negative inputs. + """ + z = z * (-7046029254386353131) + z = (z ^ ((z >> 30) & 0x3FFFFFFFF)) * (-4658895280553007687) + z = (z ^ ((z >> 27) & 0x1FFFFFFFFF)) * (-7723592293110705685) + return z ^ ((z >> 31) & 0x1FFFFFFFF) + + +def _wrap64(value: int) -> int: + """Reduce a Python int into the signed 64-bit range, wrapping. + + A salt is built on the host and handed to torch as a scalar, and torch + refuses a scalar outside the tensor's dtype. Wrapping here is what makes + "the fold wraps" true at the boundary as well as inside it. + """ + value &= (1 << 64) - 1 + return value - (1 << 64) if value >= 1 << 63 else value + + +class HoleKeyBuilder(nn.Module): + """Folds every projected slot's input values into one key per hole. + + Every member value that produces a hole contributes, discriminated by + slot, by member and by its index inside the hole. Without the first, two + slots holding the same id would match; without the last two, a two-member + slot with values ``(a, b)`` would match one with ``(b, a)`` and a permuted + multi-value item would match itself reordered -- all plausible, all wrong, + and all silent. + + Args: + prompt_plan: the compiled plan; its ``projected_slots`` fix the hole + order. + """ + + # odd multipliers, written as signed int64 so torch takes them verbatim; + # TorchScript resolves a Final class attribute as a constant + C_SLOT: Final[int] = -7046029254386353131 + C_VALUE: Final[int] = -49064778989728563 + C_INDEX: Final[int] = -2960836687051489901 + # member index and position within a hole packed into one integer, wide + # enough that a position cannot carry into the member index + MEMBER_STRIDE: Final[int] = 1 << 32 + + names: List[str] + member_names: List[List[str]] + is_sequences: List[bool] + salts: List[int] + + def __init__(self, prompt_plan: PromptPlan) -> None: + super().__init__() + self.names = [] + self.member_names = [] + self.is_sequences = [] + self.salts = [] + for seg in prompt_plan.projected_slots: + self.names.append(seg.name) + self.member_names.append(list(seg.feature_names)) + self.is_sequences.append(seg.group_type == FeatureGroupType.JAGGED_SEQUENCE) + self.salts.append(_wrap64(self.C_SLOT * int(seg.slot_id))) + self.num_slots = len(self.names) + + def _lengths(self, batch: Dict[str, torch.Tensor], member: str) -> torch.Tensor: + """Per-row item count; a dense member has one row per sample and no lengths.""" + key = member + ".lengths" + if key in batch: + return batch[key].to(torch.int64) + return torch.ones( + batch[member + ".values"].size(0), + dtype=torch.int64, + device=batch_device(batch), + ) + + def _fold_slot(self, batch: Dict[str, torch.Tensor], index: int) -> torch.Tensor: + """One projected occurrence's keys, one per hole, in sample order.""" + salt = self.salts[index] + name = self.names[index] + members = self.member_names[index] + is_sequence = self.is_sequences[index] + # the assembler's hole count: one per item of a sequence slot, one per + # sample of a DEEP slot + first = self._lengths(batch, members[0]) + num_holes = int(torch.sum(first)) if is_sequence else int(first.numel()) + keys = torch.zeros(num_holes, dtype=torch.int64, device=first.device) + for member_index in range(len(members)): + member = members[member_index] + raw = batch[member + ".values"] + key_length_key = member + ".key_lengths" + + if not is_sequence: + # a dense member contributes its float32 bit pattern verbatim, + # which is the parsed input and not a computed reduction, so + # it is stable for a given request + if raw.is_floating_point(): + width = raw.size(1) + values = ( + raw.to(torch.float32) + .contiguous() + .view(torch.int32) + .to(torch.int64) + .reshape(-1) + & 0xFFFFFFFF + ) + hole = torch.repeat_interleave( + torch.arange(num_holes, dtype=torch.int64, device=raw.device), + torch.full( + (num_holes,), width, dtype=torch.int64, device=raw.device + ), + ) + local = ( + torch.arange(width, dtype=torch.int64, device=raw.device) + .unsqueeze(0) + .expand(num_holes, width) + .reshape(-1) + ) + else: + lengths = self._lengths(batch, member) + values = raw.to(torch.int64).reshape(-1) + hole = _row_ids(lengths) + local = _within_row_index(lengths) + else: + if raw.is_floating_point(): + raise ValueError( + "prompt slot [" + + name + + "] member [" + + member + + "] is a dense sequence; the fold has no per-item boundary " + + "for it." + ) + values = raw.to(torch.int64).reshape(-1) + if key_length_key in batch: + key_lengths = batch[key_length_key].to(torch.int64).reshape(-1) + hole = _row_ids(key_lengths) + local = _within_row_index(key_lengths) + else: + hole = torch.arange( + values.numel(), dtype=torch.int64, device=values.device + ) + local = torch.zeros_like(hole) + + local = local + member_index * self.MEMBER_STRIDE + mixed = mix64(values * self.C_VALUE + local * self.C_INDEX + salt) + keys.index_add_(0, hole, mixed) + return keys + + def forward(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor: + """Fold one batch. + + Args: + batch: the parsed feature dict, keyed ``{feature}.values`` / + ``.lengths`` / ``.key_lengths`` as the data parser emits it. + + Returns: + ``(total_holes,)`` int64, row-aligned with ``hole_positions``. + """ + # an empty stream first, so a plan without a projected slot still has + # something to join + parts: List[torch.Tensor] = [ + torch.zeros(0, dtype=torch.int64, device=batch_device(batch)) + ] + for i in range(self.num_slots): + parts.append(self._fold_slot(batch, i)) + return torch.cat(parts, dim=0) diff --git a/tzrec/prompt/hole_keys_test.py b/tzrec/prompt/hole_keys_test.py new file mode 100644 index 00000000..6e4752ac --- /dev/null +++ b/tzrec/prompt/hole_keys_test.py @@ -0,0 +1,275 @@ +# 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 unittest + +import numpy as np +import torch +from torch import nn + +from tzrec.prompt.assembler import HOLE_SLOT_COUNTS, PromptAssembler +from tzrec.prompt.hole_keys import PROMPT_HOLE_KEYS, HoleKeyBuilder, mix64 +from tzrec.prompt.types import ( + FillMode, + PromptPlan, + ResolvedSidSpace, + SlotSeg, + Static, + Width, + WidthKind, +) +from tzrec.protos.model_pb2 import FeatureGroupType +from tzrec.utils.fx_util import symbolic_trace + + +def _slot( + name, + feature_names=None, + group_type=FeatureGroupType.JAGGED_SEQUENCE, + slot_id=0, +) -> SlotSeg: + """A PROJECTED slot over an ad-hoc plan.""" + return SlotSeg( + slot_id=slot_id, + name=name, + feature_names=tuple(feature_names) if feature_names is not None else (name,), + group_type=group_type, + output_key=".sequence" + if group_type == FeatureGroupType.JAGGED_SEQUENCE + else "", + fill=FillMode.PROJECTED, + width=Width(WidthKind.BOUNDED, 30), + ) + + +def _plan(segments) -> PromptPlan: + return PromptPlan( + segments=tuple(segments), + response_segments=(), + max_length=0, + max_total_length=None, + max_holes=0, + logits_suffix_len=None, + static_prefix_len=0, + projected_slots=tuple(s for s in segments if isinstance(s, SlotSeg)), + ) + + +def _tensors(raw): + return {key: torch.as_tensor(np.asarray(value)) for key, value in raw.items()} + + +class MixTest(unittest.TestCase): + def test_mix64_matches_a_host_reference(self) -> None: + """The mixer's masked shifts must reproduce SplitMix64 on negatives.""" + + def reference(value: int) -> int: + mask = (1 << 64) - 1 + z = (value * 0x9E3779B97F4A7C15) & mask + z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & mask + z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & mask + z = z ^ (z >> 31) + return z - (1 << 64) if z >= 1 << 63 else z + + values = [0, 1, -1, 2**31, -(2**31), 123456789, -987654321] + got = mix64(torch.tensor(values, dtype=torch.int64)).tolist() + self.assertEqual(got, [reference(v) for v in values]) + + +class HoleKeyBuilderTest(unittest.TestCase): + def setUp(self) -> None: + self.plan = _plan((_slot("beh"),)) + self.module = HoleKeyBuilder(self.plan) + + def _keys(self, values, lengths): + return self.module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array(lengths, dtype=np.int64), + } + ) + ) + + def test_equal_content_folds_equal(self) -> None: + """A key is a function of the hole's inputs and nothing else.""" + self.assertTrue( + torch.equal(self._keys([5, 6, 7], [3]), self._keys([5, 6, 7], [3])) + ) + + def test_different_content_folds_apart(self) -> None: + """The whole point: a changed input must not reuse the cached KV.""" + first = self._keys([5, 6, 7], [3]) + second = self._keys([5, 6, 8], [3]) + self.assertEqual(first.dtype, torch.int64) + self.assertEqual(first[:2].tolist(), second[:2].tolist()) + self.assertNotEqual(int(first[2]), int(second[2])) + + def test_keys_follow_the_assemblers_hole_order(self) -> None: + """Projected occurrence first, then sample, like ``hole_positions``.""" + a, b = _slot("a", slot_id=0), _slot("b", slot_id=1) + batch = _tensors( + { + "a.values": np.array([1, 2, 3], dtype=np.int64), + "a.lengths": np.array([1, 2], dtype=np.int64), + "b.values": np.array([4, 5, 6], dtype=np.int64), + "b.lengths": np.array([2, 1], dtype=np.int64), + } + ) + keys = HoleKeyBuilder(_plan((a, Static((7,)), b)))(batch) + counts = PromptAssembler(_plan((a, Static((7,)), b)), _SID_SPACE)(batch)[ + HOLE_SLOT_COUNTS + ] + self.assertEqual(counts.tolist(), [3, 3]) + self.assertEqual(keys.numel(), 6) + self.assertTrue(torch.equal(keys[:3], HoleKeyBuilder(_plan((a,)))(batch))) + self.assertTrue(torch.equal(keys[3:], HoleKeyBuilder(_plan((b,)))(batch))) + + def test_two_slots_holding_the_same_id_do_not_collide(self) -> None: + """The slot id salts the key, so the same value in another slot differs.""" + batch = _tensors( + { + "beh.values": np.array([5], dtype=np.int64), + "beh.lengths": np.array([1], dtype=np.int64), + } + ) + other = HoleKeyBuilder(_plan((_slot("beh", slot_id=1),))) + self.assertFalse(torch.equal(self.module(batch), other(batch))) + + def test_a_permuted_multi_value_item_does_not_collide(self) -> None: + """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" + + def keys(values): + return self.module( + _tensors( + { + "beh.values": np.array(values, dtype=np.int64), + "beh.lengths": np.array([1], dtype=np.int64), + "beh.key_lengths": np.array([3], dtype=np.int64), + } + ) + ) + + self.assertNotEqual(keys([1, 5, 9]).tolist(), keys([9, 5, 1]).tolist()) + + def test_two_members_exchanging_values_do_not_collide(self) -> None: + """Without the member index a two-member slot is order-blind.""" + module = HoleKeyBuilder(_plan((_slot("pair", feature_names=("a", "b")),))) + + def keys(first, second): + return module( + _tensors( + { + "a.values": np.array(first, dtype=np.int64), + "a.lengths": np.array([1], dtype=np.int64), + "b.values": np.array(second, dtype=np.int64), + "b.lengths": np.array([1], dtype=np.int64), + } + ) + ) + + self.assertNotEqual(keys([3], [9]).tolist(), keys([9], [3]).tolist()) + + def test_a_dense_member_folds_its_bit_pattern(self) -> None: + """A float member contributes the parsed input verbatim, so it is stable.""" + module = HoleKeyBuilder( + _plan((_slot("vec", group_type=FeatureGroupType.DEEP),)) + ) + + def keys(rows): + return module({"vec.values": torch.tensor(rows, dtype=torch.float32)}) + + self.assertEqual(keys([[0.5, 1.0]]).tolist(), keys([[0.5, 1.0]]).tolist()) + self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) + self.assertEqual(keys([[0.5, 1.0], [0.5, 1.0]]).numel(), 2) + + def test_a_dense_sequence_member_is_rejected(self) -> None: + with self.assertRaisesRegex(ValueError, "no per-item boundary"): + self.module( + { + "beh.values": torch.tensor([[0.5], [1.0]]), + "beh.lengths": torch.tensor([2]), + } + ) + + def test_no_projected_slot_folds_nothing(self) -> None: + keys = HoleKeyBuilder(_plan((Static((7,)),)))({"batch_size": torch.tensor(2)}) + self.assertEqual(keys.numel(), 0) + self.assertEqual(keys.dtype, torch.int64) + self.assertEqual(PROMPT_HOLE_KEYS, "prompt_hole_keys") + + def test_scripting_preserves_the_keys(self) -> None: + """The exported front-end folds exactly what the eager module does.""" + batch = _tensors( + { + "beh.values": np.array([7, 8, 9, 21, 22], dtype=np.int64), + "beh.lengths": np.array([3, 2], dtype=np.int64), + } + ) + scripted = torch.jit.script(self.module) + self.assertTrue(torch.equal(scripted(batch), self.module(batch))) + + def test_is_an_fx_leaf(self) -> None: + """Export traces the wrapper first; the fold must stay one opaque node.""" + + class Wrapper(nn.Module): + def __init__(self, builder): + super().__init__() + self.builder = builder + + def forward(self, batch): + return self.builder(batch) + + batch = _tensors( + { + "beh.values": np.array([7, 8, 9], dtype=np.int64), + "beh.lengths": np.array([3], dtype=np.int64), + } + ) + traced = symbolic_trace(Wrapper(self.module)) + self.assertTrue( + any( + n.op == "call_module" and n.target == "builder" + for n in traced.graph.nodes + ) + ) + self.assertTrue(torch.equal(traced(batch), self.module(batch))) + + @unittest.skipIf(not torch.cuda.is_available(), "no GPU") + def test_the_fold_is_bit_identical_across_devices(self) -> None: + """Integer addition cannot depend on the order a device reduces in.""" + batch = _tensors( + { + "beh.values": np.arange(64, dtype=np.int64), + "beh.lengths": np.array([32, 32], dtype=np.int64), + } + ) + on_cpu = self.module(batch) + on_gpu = self.module({k: v.cuda() for k, v in batch.items()}) + self.assertTrue(torch.equal(on_cpu, on_gpu.cpu())) + + +_SID_SPACE = ResolvedSidSpace( + codebook=(4, 4, 4), + num_levels=3, + base_vocab_size=1000, + level_offsets=(0, 4, 8), + band_lo=(1000, 1004, 1008), + band_hi=(1003, 1007, 1011), + target_vocab_size=1152, + sentinel_token_id=1099, + eos_token_id=2, + pad_token_id=3, +) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py index f2009dfb..d4069547 100644 --- a/tzrec/prompt/persist.py +++ b/tzrec/prompt/persist.py @@ -48,16 +48,9 @@ def write_serving_contract(compiled_prompt: CompiledPrompt, export_dir: str) -> out = os.path.join(export_dir, PROMPT_DIR) os.makedirs(out, exist_ok=True) path = os.path.join(out, PROMPT_CONTRACT_FILENAME) - sid_space = compiled_prompt.sid_space with open(path, "w") as f: json.dump( - { - "sid_space": dataclasses.asdict(sid_space) - if sid_space is not None - else None - }, - f, - indent=2, + {"sid_space": dataclasses.asdict(compiled_prompt.sid_space)}, f, indent=2 ) return path diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index e61b1486..72e0bbe7 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -16,7 +16,7 @@ physical dimension: the model resolves those at ``__init__``. """ -from dataclasses import dataclass, field +from dataclasses import dataclass from enum import Enum from typing import Mapping, Optional, Tuple, Union @@ -140,36 +140,6 @@ class SlotSeg: Segment = Union[Static, SlotSeg] -@dataclass(frozen=True) -class FoldConstants: - """Odd multipliers mixed into ``hole_keys``. - - Written as signed int64 so torch takes them verbatim: the fold wraps, and a - host that had to convert them would be a second place to get it wrong. - - They live in the plan rather than in the host so the artifact, not the - machine that runs it, decides the keys. Each closes a collision that would - otherwise produce a correct-looking prefix-cache hit: two slots holding the - same id, two members of one slot exchanging values, or a multi-value item - permuted. - - Args: - slot: multiplies ``slot_id`` into the per-hole salt. - plan: multiplies the low 64 bits of ``plan_hash``, so keys are - artifact-specific and a rolling upgrade cannot cross-match. - value: multiplies each contributing value. - index: multiplies the member-and-position index within a hole. - position: multiplies the hole index in the per-item outer fold, without - which a permuted history collides. - """ - - slot: int = -7046029254386353131 - plan: int = -4417276706812531889 - value: int = -49064778989728563 - index: int = -2960836687051489901 - position: int = -6752110988234923001 - - @dataclass(frozen=True) class PromptPlan: """The walk order the assembler follows, plus the ceilings derived from it. @@ -185,7 +155,6 @@ class PromptPlan: projected_slots: PROJECTED occurrences in emission order, which is also ascending hole position; nothing may reorder them by slot id or by shared module, because the serving scatter is positional. - fold: the constants ``hole_keys`` mixes in. """ segments: Tuple[Segment, ...] @@ -196,7 +165,6 @@ class PromptPlan: logits_suffix_len: Optional[int] static_prefix_len: int projected_slots: Tuple[SlotSeg, ...] - fold: FoldConstants = field(default_factory=FoldConstants) @dataclass(frozen=True) @@ -229,7 +197,7 @@ class CompiledPrompt: plan_hash: over all four parts; warns on mismatch. """ - sid_space: Optional[ResolvedSidSpace] + sid_space: ResolvedSidSpace prompt_plan: PromptPlan projection_plan: ProjectionPlan vocab_hash: str diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index e9569a7c..2b30b920 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -27,8 +27,7 @@ import torch from safetensors.torch import load_file -from tzrec.prompt.assembler import mix64 -from tzrec.prompt.types import FoldConstants +from tzrec.prompt.hole_keys import mix64 from tzrec.tests.prompt_test_util import export_tiny_genrec from tzrec.utils.test_util import make_test_dir @@ -127,8 +126,7 @@ def test_front_end_output_feeds_the_processor(self) -> None: self.assertEqual(out["hole_keys"].dtype, torch.int64) def test_host_item_hash_refolds_hole_keys(self) -> None: - """The per-item outer fold on the host uses the plan's position constant.""" - self.assertEqual(FoldConstants().position, _C_POSITION) + """The host's per-item outer fold reuses the front-end's mixer.""" for value in (0, 1, -1, 123456789, -(2**40)): signed = _mix64_host(value & _MASK64) signed = signed - (1 << 64) if signed >= 1 << 63 else signed diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index 32b1bfc9..f5bdc2eb 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -23,11 +23,11 @@ from tzrec.models.model import TrainWrapper from tzrec.prompt.assembler import ( CU_SEQLENS, - HOLE_KEYS, HOLE_POSITIONS, INPUT_IDS, PromptAssembler, ) +from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.tests.prompt_test_util import ( GenrecModelTestBase, assemble_into, @@ -113,21 +113,23 @@ def test_export_writes_one_serving_directory(self) -> None: ) self.assertFalse(os.path.exists(os.path.join(exported.export_dir, "sparse"))) - # the exported artifact and the collator are two call sites of one walk + # the exported artifact and the collator are two call sites of one walk; + # only the artifact folds hole keys data = exported.sample_data compiled = exported.compiled_prompt collator = PromptAssembler( - compiled.prompt_plan, - compiled.sid_space, - plan_hash=compiled.plan_hash, - include_response=False, + compiled.prompt_plan, compiled.sid_space, include_response=False )(data) front_end = torch.jit.load( os.path.join(exported.export_dir, "scripted_model.pt") ) out = front_end(data, torch.device("cpu")) - for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, HOLE_KEYS): + for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS): self.assertTrue(torch.equal(out[key], collator[key]), key) + self.assertTrue( + torch.equal(out[HOLE_KEYS], HoleKeyBuilder(compiled.prompt_plan)(data)) + ) + self.assertNotIn(HOLE_KEYS, collator) self.assertEqual( tuple(out["slot_embeds"].shape), (int(out[HOLE_POSITIONS].numel()), 32) ) diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index d4ff622a..e3f53379 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -100,11 +100,9 @@ def assemble_into( The assembled streams keyed for ``additional_infos``. """ batch = {k: torch.as_tensor(np.asarray(v)) for k, v in parsed_features.items()} - streams = PromptAssembler( - compiled_prompt.prompt_plan, - compiled_prompt.sid_space, - plan_hash=compiled_prompt.plan_hash, - )(batch) + streams = PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( + batch + ) return {PROMPT_INFO_PREFIX + k: v for k, v in streams.items()} diff --git a/tzrec/utils/fx_util.py b/tzrec/utils/fx_util.py index 01099e54..d0198622 100644 --- a/tzrec/utils/fx_util.py +++ b/tzrec/utils/fx_util.py @@ -18,7 +18,7 @@ # Modules whose forward FX cannot record -- they branch on tensor values or # turn them into Python ints -- so tracing keeps them opaque and TorchScript # compiles them whole. Matched by class name. -UNTRACEABLE_MODULES = ["ComputeJTDictToKJT", "PromptAssembler"] +UNTRACEABLE_MODULES = ["ComputeJTDictToKJT", "PromptAssembler", "HoleKeyBuilder"] def symbolic_trace( From 3a261c79d330bbc7bd728858ad74763f6ffd0145 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Mon, 7 Sep 2026 20:06:36 +0800 Subject: [PATCH 06/21] [refactor] drop the prompt digests and the restore guard Nothing in the export or in serving reads vocab_hash or plan_hash. The restore guard refused routine warm starts such as a bundle refresh, and the plan salt covered only the holes of a cache the engine namespaces anyway. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/main.py | 7 - tzrec/models/genrec_model.py | 7 - tzrec/prompt/assembler_test.py | 2 - tzrec/prompt/compile.py | 26 +--- tzrec/prompt/persist.py | 80 ++--------- tzrec/prompt/persist_test.py | 176 ------------------------- tzrec/prompt/types.py | 4 - tzrec/tests/prompt_integration_test.py | 11 -- tzrec/utils/checkpoint_util.py | 3 +- tzrec/utils/hf_export_util.py | 3 - 10 files changed, 11 insertions(+), 308 deletions(-) delete mode 100644 tzrec/prompt/persist_test.py diff --git a/tzrec/main.py b/tzrec/main.py index c469e173..9dd98442 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -74,7 +74,6 @@ from tzrec.optim.lr_scheduler import BaseLR from tzrec.optim.optimizer import TZRecOptimizer from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.persist import check_prompt_assets from tzrec.prompt.types import CompiledPrompt from tzrec.protos.data_pb2 import DataConfig, DatasetType from tzrec.protos.eval_pb2 import EvalConfig @@ -836,8 +835,6 @@ def train_and_evaluate( # Restore dataloader state before create_dataloader starts its workers dataloader_state: Optional[Dict[str, Any]] = None - if ckpt_path: - check_prompt_assets(compiled_prompt, ckpt_path) if ckpt_path and continue_train: dataloader_state = ckpt_manager.restore_dataloader_state(ckpt_path) if dataloader_state and not restore_from_model_dir: @@ -1122,7 +1119,6 @@ def evaluate( ) if checkpoint_path: - check_prompt_assets(compiled_prompt, checkpoint_path) ckpt_manager.restore( checkpoint_path, model, @@ -1213,8 +1209,6 @@ def export( features = _create_features(list(pipeline_config.feature_configs), data_config) compiled_prompt = _compile_prompt(pipeline_config, features) - if checkpoint_path: - check_prompt_assets(compiled_prompt, checkpoint_path) # Build model model = _create_model( @@ -1758,7 +1752,6 @@ def predict_checkpoint( model.eval() if checkpoint_path: - check_prompt_assets(compiled_prompt, checkpoint_path) ckpt_manager.restore( checkpoint_path, model, diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 0231bcbd..5e297660 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -306,13 +306,6 @@ def update_train_metric( """ return - def prompt_digests(self) -> Dict[str, str]: - """The contract digests the checkpoint records, for restore checking.""" - return { - "vocab_hash": self._prompt.vocab_hash, - "plan_hash": self._prompt.plan_hash, - } - def init_from_pretrained(self) -> None: """Load HF weights once, on a cold start only.""" source = self._model_config.hf_model_name_or_path diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index f2f6405b..02f1be84 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -295,8 +295,6 @@ def test_column_shaped_values_are_flattened(self) -> None: sid_space=_sid_space(), prompt_plan=plan, projection_plan=ProjectionPlan(projections={}, slot_to_module={}), - vocab_hash="v", - plan_hash="p", ) parsed = { "hist.values": np.array([[1], [6], [11], [0], [4], [8]]), diff --git a/tzrec/prompt/compile.py b/tzrec/prompt/compile.py index e9244b5e..c4611c82 100644 --- a/tzrec/prompt/compile.py +++ b/tzrec/prompt/compile.py @@ -16,11 +16,10 @@ ``__init__`` from ``group_total_dim``. """ -import hashlib import json import os import re -from typing import Any, Dict, List, Optional, Sequence, Tuple +from typing import Dict, List, Optional, Sequence, Tuple from tokenizers import Tokenizer @@ -257,14 +256,6 @@ def _special_id(tok: Tokenizer, candidates: Sequence[str]) -> int: ) -def _hash(*parts: Any) -> str: - """Stable sha256 over the given parts.""" - digest = hashlib.sha256() - for part in parts: - digest.update(repr(part).encode("utf-8")) - return digest.hexdigest() - - def compile_prompt( cfg: PromptConfig, features: Sequence[BaseFeature], @@ -362,7 +353,6 @@ def compile_prompt( if tokenizer_dir: _save_tokenizer_dir(tok, sid_space, tokenizer_dir) - tokenizer_json = tok.to_str() slot_ids = {name: i for i, name in enumerate(resolved_slots_by_name)} segs: Dict[str, SlotSeg] = {} for name, slot in resolved_slots_by_name.items(): @@ -407,19 +397,7 @@ def compile_prompt( _validate(plan) return CompiledPrompt( - sid_space=sid_space, - prompt_plan=plan, - projection_plan=projection_plan, - vocab_hash=_hash(sid_space, tokenizer_json), - plan_hash=_hash( - sid_space, - plan, - # items(), not the bare dict: iterating one yields only its keys - sorted(projection_plan.projections.items()), - # routing too: matching bodies can still be wired to other slots - sorted(projection_plan.slot_to_module.items()), - tokenizer_json, - ), + sid_space=sid_space, prompt_plan=plan, projection_plan=projection_plan ) diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py index d4069547..66a77868 100644 --- a/tzrec/prompt/persist.py +++ b/tzrec/prompt/persist.py @@ -9,26 +9,21 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Checks a checkpoint's prompt contract on restore, and publishes the SID space. - -The digests ride in the HF export metadata the checkpoint already carries, so -there is no second file that can drift from the weights beside it. Export -additionally writes ``prompt/prompt.json`` with the resolved SID space: the -token base, the per-level bands and the bundle identity a serving side needs -to build its constraint index and to refuse one built from another bundle. -The plan itself is deliberately not published -- it reaches serving compiled -into the front-end, and a copy a runtime could interpret would invite the -second assembler this design exists to prevent. +"""Publishes the SID space beside an export. + +Export writes ``prompt/prompt.json`` with the resolved SID space: the token +base, the per-level bands and the bundle identity a serving side needs to +build its constraint index and to refuse one built from another bundle. The +plan itself is deliberately not published -- it reaches serving compiled into +the front-end, and a copy a runtime could interpret would invite the second +assembler this design exists to prevent. """ import dataclasses import json import os -from typing import Dict, Optional -from tzrec.constant import HF_EXPORT_META_FILENAME from tzrec.prompt.types import CompiledPrompt -from tzrec.utils.logging_util import logger PROMPT_DIR = "prompt" PROMPT_CONTRACT_FILENAME = "prompt.json" @@ -53,62 +48,3 @@ def write_serving_contract(compiled_prompt: CompiledPrompt, export_dir: str) -> {"sid_space": dataclasses.asdict(compiled_prompt.sid_space)}, f, indent=2 ) return path - - -def read_prompt_digests(source_dir: str) -> Optional[Dict[str, str]]: - """Read the digests a checkpoint recorded, or None when it has none.""" - path = os.path.join(source_dir, HF_EXPORT_META_FILENAME) - if not os.path.exists(path): - return None - with open(path, "r") as f: - recorded = json.load(f) - if "vocab_hash" not in recorded: - return None - return recorded - - -def check_prompt_assets( - compiled_prompt: Optional[CompiledPrompt], ckpt_dir: str -) -> None: - """Compare a compiled prompt against what a checkpoint recorded. - - A ``vocab_hash`` mismatch is fatal: the decode bands would point at token - ranges the weights never learned, which produces plausible output rather - than an error. A ``plan_hash`` mismatch only reshapes the prompt, so it - warns. Absent digests are fatal too -- restoring unchecked is the one case - the guard exists to prevent. - - Args: - compiled_prompt: the freshly compiled prompt, or None when the pipeline - declares no prompt_config. - ckpt_dir: the checkpoint being restored. - - Raises: - ValueError: if the checkpoint records no digests, or its ``vocab_hash`` - disagrees with the compiled prompt. - """ - if compiled_prompt is None: - return - recorded = read_prompt_digests(ckpt_dir) - if recorded is None: - raise ValueError( - f"checkpoint [{ckpt_dir}] records no prompt digests, so its " - f"vocabulary cannot be checked against the current prompt_config. " - f"Restoring unchecked risks decode bands that address rows these " - f"weights never learned, so this is fatal rather than a warning." - ) - - if recorded.get("vocab_hash") != compiled_prompt.vocab_hash: - raise ValueError( - f"prompt vocabulary does not match checkpoint [{ckpt_dir}]: the " - f"checkpoint was trained against {recorded.get('vocab_hash')} but " - f"prompt_config now compiles to {compiled_prompt.vocab_hash}. The " - f"SID space or the tokenizer changed, so the decode bands no longer " - f"address the rows these weights learned." - ) - if recorded.get("plan_hash") != compiled_prompt.plan_hash: - logger.warning( - f"prompt plan differs from checkpoint [{ckpt_dir}]: the vocabulary " - f"matches, so the weights are usable, but the template, slots or " - f"projections changed." - ) diff --git a/tzrec/prompt/persist_test.py b/tzrec/prompt/persist_test.py deleted file mode 100644 index 17c92669..00000000 --- a/tzrec/prompt/persist_test.py +++ /dev/null @@ -1,176 +0,0 @@ -# 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 json -import os -import unittest -from unittest import mock - -from tzrec.constant import HF_EXPORT_META_FILENAME -from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.persist import ( - check_prompt_assets, - read_prompt_digests, -) -from tzrec.protos.prompt_pb2 import PromptConfig -from tzrec.tests.prompt_test_util import ( - create_prompt_feature, - create_prompt_tokenizer, -) -from tzrec.utils.test_util import make_test_dir - -_WORDS = ["History", "Predict", ":", "", "<|im_end|>"] - - -def _record(compiled_prompt, ckpt_dir: str) -> None: - """Write the digests where write_hf_assets puts them.""" - os.makedirs(ckpt_dir, exist_ok=True) - with open(os.path.join(ckpt_dir, HF_EXPORT_META_FILENAME), "w") as f: - json.dump( - { - "backbone_state_dict_prefix": "lm.", - "vocab_hash": compiled_prompt.vocab_hash, - "plan_hash": compiled_prompt.plan_hash, - }, - f, - ) - - -class PromptPersistTest(unittest.TestCase): - def setUp(self) -> None: - self.test_dir = make_test_dir() - self.tok_path = create_prompt_tokenizer( - os.path.join(self.test_dir, "tok.json"), _WORDS - ) - self.features = [ - create_prompt_feature( - 'sequence_raw_feature { feature_name: "hist" expression: "user:hist" }' - ), - ] - - def _compile(self, codebook=(4, 4, 4), prompt="History : {{hist}}"): - cfg = PromptConfig( - tokenizer_path=self.tok_path, prompt=prompt, response="{{answer}}" - ) - cfg.sid_space.codebook.extend(codebook) - return compile_prompt(cfg, self.features, ["answer"]) - - def test_a_changed_codebook_is_fatal(self) -> None: - ckpt = os.path.join(self.test_dir, "model.ckpt-1") - _record(self._compile(codebook=(4, 4, 4)), ckpt) - with self.assertRaisesRegex(ValueError, "does not match checkpoint"): - check_prompt_assets(self._compile(codebook=(8, 8, 8)), ckpt) - - def test_a_changed_projection_body_warns(self) -> None: - # a body change must not hide behind an unchanged projection name - def compile_with(hidden_units): - cfg = PromptConfig( - tokenizer_path=self.tok_path, - prompt="History : {{hist}} {{prof}}", - response="{{answer}}", - ) - cfg.sid_space.codebook.extend((4, 4, 4)) - slot = cfg.slots.add(name="prof") - slot.feature_names.append("prof") - slot.projection.mlp.hidden_units.extend(hidden_units) - features = self.features + [ - create_prompt_feature( - 'sequence_id_feature { feature_name: "prof" ' - 'expression: "user:prof" num_buckets: 16 embedding_dim: 8 ' - "sequence_length: 2 }" - ) - ] - return compile_prompt(cfg, features, ["answer"]) - - ckpt = os.path.join(self.test_dir, "model.ckpt-proj") - _record(compile_with([16]), ckpt) - widened = compile_with([256, 128]) - # only the projection changed, so the vocabulary is still usable - self.assertEqual(widened.vocab_hash, read_prompt_digests(ckpt)["vocab_hash"]) - self.assertNotEqual(widened.plan_hash, read_prompt_digests(ckpt)["plan_hash"]) - with mock.patch("tzrec.prompt.persist.logger.warning") as warning: - check_prompt_assets(widened, ckpt) - warning.assert_called_once() - - def test_swapped_projection_routing_warns(self) -> None: - # identical bodies, so only slot_to_module differs - def compile_with(pa_module, pb_module): - cfg = PromptConfig( - tokenizer_path=self.tok_path, - prompt="History : {{hist}} {{pa}} {{pb}}", - response="{{answer}}", - ) - cfg.sid_space.codebook.extend((4, 4, 4)) - features = list(self.features) - for name, module_id in (("pa", pa_module), ("pb", pb_module)): - slot = cfg.slots.add(name=name, projection_name=module_id) - slot.feature_names.append(name) - slot.projection.mlp.hidden_units.extend([16]) - features.append( - create_prompt_feature( - f'sequence_id_feature {{ feature_name: "{name}" ' - f'expression: "user:{name}" num_buckets: 16 ' - "embedding_dim: 8 sequence_length: 2 }" - ) - ) - return compile_prompt(cfg, features, ["answer"]) - - ckpt = os.path.join(self.test_dir, "model.ckpt-route") - _record(compile_with("X", "Y"), ckpt) - swapped = compile_with("Y", "X") - self.assertEqual(swapped.vocab_hash, read_prompt_digests(ckpt)["vocab_hash"]) - self.assertNotEqual(swapped.plan_hash, read_prompt_digests(ckpt)["plan_hash"]) - with mock.patch("tzrec.prompt.persist.logger.warning") as warning: - check_prompt_assets(swapped, ckpt) - warning.assert_called_once() - - def test_a_changed_template_only_warns(self) -> None: - ckpt = os.path.join(self.test_dir, "model.ckpt-1") - _record(self._compile(), ckpt) - changed_compiled_prompt = self._compile(prompt="Predict : {{hist}}") - # the vocabulary is untouched, so the weights are still usable - self.assertEqual( - changed_compiled_prompt.vocab_hash, - read_prompt_digests(ckpt)["vocab_hash"], - ) - self.assertNotEqual( - changed_compiled_prompt.plan_hash, - read_prompt_digests(ckpt)["plan_hash"], - ) - with mock.patch("tzrec.prompt.persist.logger.warning") as warning: - check_prompt_assets(changed_compiled_prompt, ckpt) - warning.assert_called_once() - - def test_a_checkpoint_without_assets_is_fatal(self) -> None: - bare = os.path.join(self.test_dir, "model.ckpt-bare") - os.makedirs(bare, exist_ok=True) - self.assertIsNone(read_prompt_digests(bare)) - with self.assertRaisesRegex(ValueError, "records no prompt digests"): - check_prompt_assets(self._compile(), bare) - - def test_hf_metadata_without_digests_is_fatal(self) -> None: - ckpt = os.path.join(self.test_dir, "model.ckpt-nodigest") - os.makedirs(ckpt, exist_ok=True) - with open(os.path.join(ckpt, HF_EXPORT_META_FILENAME), "w") as f: - json.dump({"backbone_state_dict_prefix": "lm."}, f) - - self.assertIsNone(read_prompt_digests(ckpt)) - with self.assertRaisesRegex(ValueError, "records no prompt digests"): - check_prompt_assets(self._compile(), ckpt) - - def test_no_prompt_config_is_a_no_op(self) -> None: - with mock.patch("tzrec.prompt.persist.logger.warning") as warning: - check_prompt_assets(None, os.path.join(self.test_dir, "nowhere")) - warning.assert_not_called() - - -if __name__ == "__main__": - unittest.main() diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index 72e0bbe7..20b0f365 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -193,12 +193,8 @@ class CompiledPrompt: sid_space: the resolved SID token space. prompt_plan: assembler walk order and ceilings. projection_plan: projection topology. - vocab_hash: over sid_space and tokenizer.json; fatal on mismatch. - plan_hash: over all four parts; warns on mismatch. """ sid_space: ResolvedSidSpace prompt_plan: PromptPlan projection_plan: ProjectionPlan - vocab_hash: str - plan_hash: str diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index f5bdc2eb..2d25479e 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -54,17 +54,6 @@ def _batch_from_codes(self, hist, answer): batch.additional_infos.update(assemble_into(self.compiled_prompt, parsed)) return batch - def test_written_digests_satisfy_the_restore_guard(self) -> None: - from tzrec.prompt.persist import check_prompt_assets - from tzrec.utils.hf_export_util import write_hf_assets - - model = self._model() - ckpt = os.path.join(self.test_dir, "model.ckpt-1") - write_hf_assets(model, ckpt) - - check_prompt_assets(self.compiled_prompt, ckpt) - self.assertTrue(os.path.exists(os.path.join(ckpt, "hf_export_meta.json"))) - def test_model_resizes_to_target_vocab_size(self) -> None: model = self._model() rows = model.lm.get_input_embeddings().weight.shape[0] diff --git a/tzrec/utils/checkpoint_util.py b/tzrec/utils/checkpoint_util.py index bd9593a8..5481d282 100644 --- a/tzrec/utils/checkpoint_util.py +++ b/tzrec/utils/checkpoint_util.py @@ -431,8 +431,7 @@ def save( """Save a checkpoint at the given step, then request an async prune. For HF-backed models, writes the config, optional tokenizer, and the - state dict metadata and contract digests HF conversion and restore - checking read. + state dict metadata HF conversion reads. """ ckpt_dir = os.path.join(self._model_dir, f"model.ckpt-{step}") save_model(ckpt_dir, model, optimizer, dense_ema) diff --git a/tzrec/utils/hf_export_util.py b/tzrec/utils/hf_export_util.py index 44095f57..0778110d 100644 --- a/tzrec/utils/hf_export_util.py +++ b/tzrec/utils/hf_export_util.py @@ -71,9 +71,6 @@ def write_hf_assets(wrapped_model: nn.Module, save_dir: str) -> None: ) prefix = checkpoint_util._strip_dmp_prefix(raw_prefix) meta = {"backbone_state_dict_prefix": prefix + ("." if prefix else "")} - digests = getattr(inner, "prompt_digests", None) - if digests is not None: - meta.update(digests()) with open(os.path.join(save_dir, HF_EXPORT_META_FILENAME), "w") as f: json.dump(meta, f, indent=2) From ef71e5df69320a76529871106f8e997af0fbd7d1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 11:53:35 +0800 Subject: [PATCH 07/21] [refactor] rename Genrec to GenRec Each acronym keeps its own casing, as DeepFM and DlrmHSTU do, and the serving side already spells it PromptGenRecForCausalLM. Config field names and module file names are unchanged. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/main.py | 6 +++--- tzrec/models/genrec_causal_lm_model.py | 12 ++++++------ tzrec/models/genrec_causal_lm_model_test.py | 8 ++++---- tzrec/models/genrec_model.py | 16 ++++++++-------- tzrec/models/genrec_model_test.py | 18 +++++++++--------- tzrec/protos/model.proto | 2 +- tzrec/protos/models/genrec_model.proto | 6 +++--- tzrec/tests/genrec_serving_contract_test.py | 2 +- tzrec/tests/prompt_integration_test.py | 6 +++--- tzrec/tests/prompt_test_util.py | 10 +++++----- 10 files changed, 43 insertions(+), 43 deletions(-) diff --git a/tzrec/main.py b/tzrec/main.py index 9dd98442..1b4337df 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -52,7 +52,7 @@ BaseFeature, create_features, ) -from tzrec.models.genrec_model import BaseGenrecModel, GenrecFrontEnd +from tzrec.models.genrec_model import BaseGenRecModel, GenRecFrontEnd from tzrec.models.match_model import ( MatchModel, MatchTower, @@ -1264,12 +1264,12 @@ def export( os.path.join(export_dir, "model"), assets=assets, ) - elif isinstance(model.model, BaseGenrecModel): + elif isinstance(model.model, BaseGenRecModel): # tzrec serves the prompt front-end; the LM rides beside it as # HuggingFace weights for the engine that decodes export_model( ori_pipeline_config, - InferWrapper(GenrecFrontEnd(model.model)), + InferWrapper(GenRecFrontEnd(model.model)), checkpoint_path, export_dir, assets=assets, diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index c1b6bf2f..d72471c7 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -26,7 +26,7 @@ from tzrec.datasets.utils import Batch from tzrec.features.feature import BaseFeature -from tzrec.models.genrec_model import BaseGenrecModel +from tzrec.models.genrec_model import BaseGenRecModel from tzrec.modules.dynamic_beam import capped_beam_widths, dynamic_beam_search from tzrec.prompt.assembler import ( PROMPT_CU_SEQLENS, @@ -36,10 +36,10 @@ ) from tzrec.prompt.types import CompiledPrompt from tzrec.protos.model_pb2 import ModelConfig -from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -class GenrecCausalLMModel(BaseGenrecModel): +class GenRecCausalLMModel(BaseGenRecModel): """An HF causal LM driven by a compiled prompt. Args: @@ -71,7 +71,7 @@ def __init__( self._generated_sids_key = common.generated_sids_key self._read_beam_config(common) - def _read_beam_config(self, common: GenrecModelConfig) -> None: + def _read_beam_config(self, common: GenRecModelConfig) -> None: """Parse the decode knobs; the schedule must match the codebook. Args: @@ -221,7 +221,7 @@ def _left_pad_packed_inputs( @torch.fx.wrap def _fx_wrapped_forward( - model: "GenrecCausalLMModel", embeds: torch.Tensor, batch: Batch + model: "GenRecCausalLMModel", embeds: torch.Tensor, batch: Batch ) -> Tuple[torch.Tensor, torch.Tensor]: """Hide the padded forward from FX. @@ -242,7 +242,7 @@ def _fx_wrapped_forward( @torch.fx.wrap def _fx_wrapped_generate( - model: "GenrecCausalLMModel", embeds: torch.Tensor, batch: Batch + model: "GenRecCausalLMModel", embeds: torch.Tensor, batch: Batch ) -> torch.Tensor: """Hide the decode loop from FX. diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index e269fc19..b1cce726 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -16,7 +16,7 @@ from parameterized import parameterized from tzrec.datasets.utils import Batch -from tzrec.models.genrec_causal_lm_model import GenrecCausalLMModel +from tzrec.models.genrec_causal_lm_model import GenRecCausalLMModel from tzrec.prompt.assembler import ( PROMPT_CU_SEQLENS, PROMPT_INPUT_IDS, @@ -25,7 +25,7 @@ ) from tzrec.tests.prompt_test_util import ( _CODEBOOK, - GenrecModelTestBase, + GenRecModelTestBase, offset_sid_codes, ) from tzrec.utils.test_util import parameterized_name_func @@ -48,7 +48,7 @@ def test_packs_rows_of_different_lengths(self) -> None: PROMPT_RESPONSE_LENGTHS: response_lengths, } ) - model = GenrecCausalLMModel.__new__(GenrecCausalLMModel) + model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) torch.nn.Module.__init__(model) model._ignore_index = ignore @@ -72,7 +72,7 @@ def test_packs_rows_of_different_lengths(self) -> None: ) -class GenrecCausalLMModelTest(GenrecModelTestBase): +class GenRecCausalLMModelTest(GenRecModelTestBase): """The decode schedule and the training forward, both subclass-owned.""" @parameterized.expand( diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 5e297660..0c24d5fc 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -16,7 +16,7 @@ supplies the digests a checkpoint records. A family subclass owns its forward and decode path. -``GenrecFrontEnd`` is the half of the model tzrec serves: the assembled prompt +``GenRecFrontEnd`` is the half of the model tzrec serves: the assembled prompt and the projected slots, everything before the LM's embedding gather. It is exported like any tzrec model; the LM itself is handed to an LLM engine as the HuggingFace weights beside it. @@ -52,7 +52,7 @@ from tzrec.prompt.persist import PROMPT_DIR, TOKENIZER_DIR, write_serving_contract from tzrec.prompt.types import CompiledPrompt, PromptPlan from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig -from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.utils import config_util, env_util from tzrec.utils.hf_export_util import dcp_to_hf, write_composite_config @@ -61,9 +61,9 @@ SLOT_EMBEDS = "slot_embeds" _PARAM_DTYPE: Dict[int, torch.dtype] = { - GenrecModelConfig.FP32: torch.float32, - GenrecModelConfig.BF16: torch.bfloat16, - GenrecModelConfig.FP16: torch.float16, + GenRecModelConfig.FP32: torch.float32, + GenRecModelConfig.BF16: torch.bfloat16, + GenRecModelConfig.FP16: torch.float16, } _REQUIRED_LM_ATTRS: Tuple[str, ...] = ( @@ -73,7 +73,7 @@ ) -class BaseGenrecModel(BaseModel): +class BaseGenRecModel(BaseModel): """An HF backbone driven by a compiled prompt. Args: @@ -346,7 +346,7 @@ def project_slots( return parts[0] if len(parts) == 1 else torch.cat(parts) -class GenrecFrontEnd(nn.Module): +class GenRecFrontEnd(nn.Module): """The served half of a genrec model, exported like any tzrec model. It shares the model's embedding group and projections, so under the @@ -364,7 +364,7 @@ class GenrecFrontEnd(nn.Module): model: the genrec model to serve. """ - def __init__(self, model: BaseGenrecModel) -> None: + def __init__(self, model: BaseGenRecModel) -> None: super().__init__() if acc_utils.is_aot() or acc_utils.is_trt() or env_util.use_rtp(): raise ValueError( diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 402490f6..440623c4 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -23,7 +23,7 @@ from tzrec.models.genrec_model import ( _PARAM_DTYPE, SLOT_EMBEDS, - GenrecFrontEnd, + GenRecFrontEnd, project_slots, ) from tzrec.models.model import ScriptWrapper @@ -36,13 +36,13 @@ ) from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder -from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig from tzrec.tests.prompt_test_util import ( _CODEBOOK, _HIST, - GenrecModelTestBase, + GenRecModelTestBase, create_prompt_feature, offset_sid_codes, projected_feature, @@ -55,7 +55,7 @@ ) -class BaseGenrecModelTest(GenrecModelTestBase): +class BaseGenRecModelTest(GenRecModelTestBase): """Shared causal-LM behavior, reached through its concrete subclass.""" def test_tokens_to_local_codes_undoes_shifts_and_groups_beams(self) -> None: @@ -144,7 +144,7 @@ def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: self.assertIsNotNone(proj.head.weight.grad) @parameterized.expand( - [[GenrecModelConfig.BF16], [GenrecModelConfig.FP16]], + [[GenRecModelConfig.BF16], [GenRecModelConfig.FP16]], name_func=parameterized_name_func, ) def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> None: @@ -218,7 +218,7 @@ def test_init_from_pretrained_replaces_the_empty_weights(self) -> None: torch.testing.assert_close(after, expected) -class GenrecFrontEndTest(GenrecModelTestBase): +class GenRecFrontEndTest(GenRecModelTestBase): """The served half of the model, under the same wrapper every export uses.""" def setUp(self) -> None: @@ -248,7 +248,7 @@ def setUp(self) -> None: "beh.values": torch.tensor([3, 9]), "beh.lengths": torch.tensor([2]), } - self.wrapped = ScriptWrapper(GenrecFrontEnd(self.model)) + self.wrapped = ScriptWrapper(GenRecFrontEnd(self.model)) def test_front_end_returns_the_walk_and_the_projected_slots(self) -> None: out = self.wrapped(self.data) @@ -294,7 +294,7 @@ def test_export_assets_write_the_engine_side(self) -> None: out_dir = os.path.join(self.test_dir, "export") os.makedirs(out_dir) shutil.copy(os.path.join(ckpt, "config.json"), out_dir) - GenrecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) + GenRecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) with open(os.path.join(out_dir, "config.json"), "r") as f: self.assertEqual(json.load(f)["model_type"], "prompt_genrec") for name in ( @@ -306,7 +306,7 @@ def test_export_assets_write_the_engine_side(self) -> None: pipeline.export_config.use_dense_ema = True with self.assertRaisesRegex(ValueError, "Dense EMA"): - GenrecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) + GenRecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) if __name__ == "__main__": diff --git a/tzrec/protos/model.proto b/tzrec/protos/model.proto index 56b41686..a7512627 100644 --- a/tzrec/protos/model.proto +++ b/tzrec/protos/model.proto @@ -89,7 +89,7 @@ message ModelConfig { SidRqvae sid_rqvae = 600; SidRqkmeans sid_rqkmeans = 601; - GenrecCausalLMModel genrec_causal_lm_model = 701; + GenRecCausalLMModel genrec_causal_lm_model = 701; CustomModel custom_model = 1000; } diff --git a/tzrec/protos/models/genrec_model.proto b/tzrec/protos/models/genrec_model.proto index dcee9766..d8d86edd 100644 --- a/tzrec/protos/models/genrec_model.proto +++ b/tzrec/protos/models/genrec_model.proto @@ -1,7 +1,7 @@ syntax = "proto2"; package tzrec.protos; -message GenrecModelConfig { +message GenRecModelConfig { // Nested so it does not collide with another model's dtype enum: proto2 // enum values are siblings of their enclosing scope. enum ParamDtype { @@ -26,8 +26,8 @@ message GenrecModelConfig { optional ParamDtype lm_parameter_dtype = 5 [default = FP32]; } -message GenrecCausalLMModel { - optional GenrecModelConfig common = 1; +message GenRecCausalLMModel { + optional GenRecModelConfig common = 1; // HF hub id or local path. Names the architecture used for every model build // and the weights loaded by init_from_pretrained at cold start; the vocabulary diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index 2b30b920..4471c5ae 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -51,7 +51,7 @@ def _item_hash(keys) -> int: return total -class GenrecServingContractTest(unittest.TestCase): +class GenRecServingContractTest(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.exported = export_tiny_genrec(make_test_dir(), bundle_uuid="bundle-test") diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index 2d25479e..4f63bad6 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -29,7 +29,7 @@ ) from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.tests.prompt_test_util import ( - GenrecModelTestBase, + GenRecModelTestBase, assemble_into, export_tiny_genrec, offset_sid_codes, @@ -40,7 +40,7 @@ _WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] -class PromptStackIntegrationTest(GenrecModelTestBase): +class PromptStackIntegrationTest(GenRecModelTestBase): """compile -> assemble -> model, on the real code path.""" def _batch_from_codes(self, hist, answer): @@ -78,7 +78,7 @@ def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(model)) -class GenrecExportIntegrationTest(unittest.TestCase): +class GenRecExportIntegrationTest(unittest.TestCase): """checkpoint -> export -> the artifacts an LLM engine and a processor load.""" def setUp(self) -> None: diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index e3f53379..cfc7635b 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -119,8 +119,8 @@ def projected_feature(name: str, dim: int) -> str: ) -class GenrecModelTestBase(unittest.TestCase): - """Builds a real GenrecCausalLMModel over a tiny backbone and prompt.""" +class GenRecModelTestBase(unittest.TestCase): + """Builds a real GenRecCausalLMModel over a tiny backbone and prompt.""" def setUp(self) -> None: """Build the tiny backbone, tokenizer and compiled prompt.""" @@ -191,7 +191,7 @@ def write_genrec_checkpoint(model: torch.nn.Module, ckpt_dir: str) -> str: @dataclasses.dataclass -class ExportedGenrec: +class ExportedGenRec: """A tiny genrec model, its checkpoint and its export. Args: @@ -252,7 +252,7 @@ def export_tiny_genrec( bundle_uuid: Optional[str] = None, use_dense_ema: bool = False, env: Optional[Dict[str, str]] = None, -) -> ExportedGenrec: +) -> ExportedGenRec: """Train nothing, save a checkpoint of a tiny genrec model, and export it. The prompt carries an INLINE SID history and, when ``projected``, one @@ -365,7 +365,7 @@ def export_tiny_genrec( for k, v in batch.to_dict(sparse_dtype=torch.int64).items() if not k.startswith(PROMPT_INFO_PREFIX) } - return ExportedGenrec( + return ExportedGenRec( config=config, config_path=config_path, features=features, From db61bbc5da354f1441de883bbd000302c1edddd7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 11:58:28 +0800 Subject: [PATCH 08/21] [refactor] write the HF assets from export, not from the front-end The served module is a model, not an exporter: export() writes the HF weights, composite config, tokenizer and prompt.json after export_model, through export_hf_assets, and the duck-typed export_assets hook is gone. The bundle identity leaves ResolvedSidSpace for a top-level prompt.json key read from the manifest at export, and the assembler scatters every segment in one index_copy_. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/main.py | 13 +++++ tzrec/models/genrec_model.py | 43 ++--------------- tzrec/models/genrec_model_test.py | 27 ----------- tzrec/prompt/assembler.py | 4 +- tzrec/prompt/compile.py | 16 ++----- tzrec/prompt/persist.py | 32 +++++++++++-- tzrec/prompt/types.py | 5 -- tzrec/tests/genrec_serving_contract_test.py | 4 +- tzrec/utils/export_util.py | 19 -------- tzrec/utils/hf_export_util.py | 53 ++++++++++++++++++++- 10 files changed, 104 insertions(+), 112 deletions(-) diff --git a/tzrec/main.py b/tzrec/main.py index 1b4337df..c4c83cd2 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -101,6 +101,7 @@ export_model, ) from tzrec.utils.filesystem_util import url_to_fs +from tzrec.utils.hf_export_util import export_hf_assets from tzrec.utils.load_class import import_class from tzrec.utils.logging_util import ProgressLogger, logger from tzrec.utils.online_dense_export_util import OnlineDenseExportManager @@ -1267,6 +1268,16 @@ def export( elif isinstance(model.model, BaseGenRecModel): # tzrec serves the prompt front-end; the LM rides beside it as # HuggingFace weights for the engine that decodes + if config_util.use_dense_ema( + pipeline_config.export_config, pipeline_config.train_config + ): + raise ValueError( + "HF export: dcp_to_hf reads /model, so it cannot " + "serve Dense EMA parameters. Set export_config.use_dense_ema to " + "false to export the raw weights." + ) + if not checkpoint_path: + raise ValueError("HF export: no checkpoint found to convert.") export_model( ori_pipeline_config, InferWrapper(GenRecFrontEnd(model.model)), @@ -1275,6 +1286,8 @@ def export( assets=assets, additional_export_config=additional_export_config, ) + if is_rank_zero: + export_hf_assets(pipeline_config, features, checkpoint_path, export_dir) else: export_model( ori_pipeline_config, diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 0c24d5fc..8d57ab99 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -12,9 +12,8 @@ """Shared causal-LM plumbing for generative recommendation models. This layer builds an empty causal LM, resizes its vocabulary, wires slot -projections, converts SID coordinate systems, scores the response window and -supplies the digests a checkpoint records. A family subclass owns its forward -and decode path. +projections, converts SID coordinate systems and scores the response window. +A family subclass owns its forward and decode path. ``GenRecFrontEnd`` is the half of the model tzrec serves: the assembled prompt and the projected slots, everything before the LM's embedding gather. It is @@ -23,7 +22,6 @@ """ import inspect -import os from typing import Any, Dict, List, Optional, Sequence, Tuple import torch @@ -47,15 +45,11 @@ PROMPT_HOLE_SLOT_COUNTS, PROMPT_INPUT_IDS, ) -from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, PROMPT_HOLE_KEYS -from tzrec.prompt.persist import PROMPT_DIR, TOKENIZER_DIR, write_serving_contract from tzrec.prompt.types import CompiledPrompt, PromptPlan from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -from tzrec.protos.pipeline_pb2 import EasyRecConfig -from tzrec.utils import config_util, env_util -from tzrec.utils.hf_export_util import dcp_to_hf, write_composite_config +from tzrec.utils import env_util from tzrec.utils.logging_util import logger SLOT_EMBEDS = "slot_embeds" @@ -424,34 +418,3 @@ def predict(self, batch: Batch) -> Dict[str, torch.Tensor]: 0, 0, dtype=torch.float32, device=infos[PROMPT_INPUT_IDS].device ) return out - - def export_assets( - self, pipeline_config: EasyRecConfig, checkpoint_path: str, save_dir: str - ) -> None: - """Write what an LLM engine reads beside the scripted front-end. - - The HuggingFace weights and composite config, the extended tokenizer and - the SID space. Called by the export on rank 0, inside its save dir. - - Args: - pipeline_config: the pipeline being exported. - checkpoint_path: the checkpoint the weights come from. - save_dir: the export directory. - """ - if config_util.use_dense_ema( - pipeline_config.export_config, pipeline_config.train_config - ): - raise ValueError( - "HF export: dcp_to_hf reads /model, so it cannot " - "serve Dense EMA parameters. Set export_config.use_dense_ema to " - "false to export the raw weights." - ) - dcp_to_hf(checkpoint_path, save_dir) - write_composite_config(save_dir) - compile_prompt( - pipeline_config.prompt_config, - self._features, - list(pipeline_config.data_config.label_fields), - tokenizer_dir=os.path.join(save_dir, PROMPT_DIR, TOKENIZER_DIR), - ) - write_serving_contract(self._prompt, save_dir) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 440623c4..11119842 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -9,9 +9,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import json -import os -import shutil import unittest import torch @@ -37,7 +34,6 @@ from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig from tzrec.tests.prompt_test_util import ( _CODEBOOK, @@ -46,7 +42,6 @@ create_prompt_feature, offset_sid_codes, projected_feature, - write_genrec_checkpoint, ) from tzrec.utils.fx_util import symbolic_trace from tzrec.utils.state_dict_util import init_parameters @@ -286,28 +281,6 @@ def test_front_end_traces_and_scripts(self) -> None: for key, value in eager.items(): self.assertTrue(torch.allclose(out[key].float(), value.float()), key) - def test_export_assets_write_the_engine_side(self) -> None: - ckpt = write_genrec_checkpoint(self.model, os.path.join(self.test_dir, "ckpt")) - pipeline = EasyRecConfig() - pipeline.prompt_config.CopyFrom(self.prompt_config) - pipeline.data_config.label_fields.append("answer") - out_dir = os.path.join(self.test_dir, "export") - os.makedirs(out_dir) - shutil.copy(os.path.join(ckpt, "config.json"), out_dir) - GenRecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) - with open(os.path.join(out_dir, "config.json"), "r") as f: - self.assertEqual(json.load(f)["model_type"], "prompt_genrec") - for name in ( - "model.safetensors", - "prompt/prompt.json", - "prompt/tokenizer/tokenizer.json", - ): - self.assertTrue(os.path.exists(os.path.join(out_dir, name)), name) - - pipeline.export_config.use_dense_ema = True - with self.assertRaisesRegex(ValueError, "Dense EMA"): - GenRecFrontEnd(self.model).export_assets(pipeline, ckpt, out_dir) - if __name__ == "__main__": unittest.main() diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index 0fb8c00f..fdae9824 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -461,14 +461,16 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: hole_parts.append(torch.zeros(0, dtype=torch.int64, device=device)) response_lengths = torch.zeros(batch_size, dtype=torch.int64, device=device) + dests: List[torch.Tensor] = [] for i in range(self.num_segments): dest = _destinations(row_start + seg_offsets[i], seg_lens[i]) - out.index_copy_(0, dest, seg_values[i]) + dests.append(dest) slot = self.hole_slots[i] if slot >= 0: hole_parts[slot + 1] = dest if i >= self.num_body: response_lengths = response_lengths + seg_lens[i] + out.index_copy_(0, torch.cat(dests, dim=0), torch.cat(seg_values, dim=0)) hole_positions = torch.cat(hole_parts, dim=0) # how many of those holes each projected occurrence owns, so a host can diff --git a/tzrec/prompt/compile.py b/tzrec/prompt/compile.py index c4611c82..f28666bc 100644 --- a/tzrec/prompt/compile.py +++ b/tzrec/prompt/compile.py @@ -133,13 +133,13 @@ def _render_sid_tokens(sid_space: SidSpace) -> List[str]: return [fmt.replace("{i}", str(i)) for i in range(sum(sid_space.codebook))] -def _read_manifest(path: str) -> Tuple[List[int], str]: - """Read ``codebook`` and the bundle identity from a SID manifest.""" +def _read_manifest(path: str) -> List[int]: + """Read ``codebook`` from a SID manifest.""" if not os.path.exists(path): raise ValueError(f"sid_space.manifest_path [{path}] does not exist.") with open(path, "r") as f: manifest = json.load(f) - return [int(c) for c in manifest["codebook"]], str(manifest.get("bundle_uuid", "")) + return [int(c) for c in manifest["codebook"]] def _build_sid_space( @@ -159,15 +159,8 @@ def _build_sid_space( if any(c <= 0 for c in codebook): raise ValueError(f"every codebook size must be positive, got {codebook}.") - bundle_uuid = "" if space.HasField("manifest_path"): - declared, bundle_uuid = _read_manifest(space.manifest_path) - if not bundle_uuid: - logger.warning( - f"the SID manifest at [{space.manifest_path}] carries no " - f"bundle_uuid, so serving cannot verify that a catalog belongs " - f"to this bundle by identity rather than by path." - ) + declared = _read_manifest(space.manifest_path) if declared != codebook: raise ValueError( f"sid_space.codebook {codebook} does not match the manifest at " @@ -219,7 +212,6 @@ def _build_sid_space( sentinel_token_id=sentinel_id, eos_token_id=_special_id(tok, ("<|im_end|>", "<|endoftext|>")), pad_token_id=_special_id(tok, ("<|endoftext|>", "<|im_end|>")), - bundle_uuid=bundle_uuid, ) diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py index 66a77868..c99a7414 100644 --- a/tzrec/prompt/persist.py +++ b/tzrec/prompt/persist.py @@ -23,18 +23,38 @@ import json import os -from tzrec.prompt.types import CompiledPrompt +from tzrec.prompt.types import ResolvedSidSpace PROMPT_DIR = "prompt" PROMPT_CONTRACT_FILENAME = "prompt.json" TOKENIZER_DIR = "tokenizer" -def write_serving_contract(compiled_prompt: CompiledPrompt, export_dir: str) -> str: - """Write ``prompt/prompt.json``: the resolved SID space. +def read_bundle_uuid(manifest_path: str) -> str: + """The identity of the SID bundle a manifest describes, empty without one. Args: - compiled_prompt: the compiled prompt. + manifest_path: ``sid_space.manifest_path``, possibly unset. + + Returns: + The manifest's ``bundle_uuid``, or ``""`` when there is no manifest or + it records none. + """ + if not manifest_path: + return "" + with open(manifest_path, "r") as f: + return str(json.load(f).get("bundle_uuid", "")) + + +def write_serving_contract( + sid_space: ResolvedSidSpace, bundle_uuid: str, export_dir: str +) -> str: + """Write ``prompt/prompt.json``: the resolved SID space and the bundle identity. + + Args: + sid_space: the resolved SID token space. + bundle_uuid: the bundle the codebook was read from, so an index builder + can refuse a catalog from another bundle; empty when unknown. export_dir: the export directory. Returns: @@ -45,6 +65,8 @@ def write_serving_contract(compiled_prompt: CompiledPrompt, export_dir: str) -> path = os.path.join(out, PROMPT_CONTRACT_FILENAME) with open(path, "w") as f: json.dump( - {"sid_space": dataclasses.asdict(compiled_prompt.sid_space)}, f, indent=2 + {"sid_space": dataclasses.asdict(sid_space), "bundle_uuid": bundle_uuid}, + f, + indent=2, ) return path diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index 20b0f365..b8d5aa0d 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -84,10 +84,6 @@ class ResolvedSidSpace: slot is projected. eos_token_id: end-of-sequence id of the extended tokenizer. pad_token_id: padding id of the extended tokenizer. - bundle_uuid: identity of the SID bundle this space was compiled - against, empty when no manifest was read. Serving refuses a - catalog whose bundle differs: a copied artifact's path proves - nothing. """ codebook: Tuple[int, ...] @@ -100,7 +96,6 @@ class ResolvedSidSpace: sentinel_token_id: Optional[int] eos_token_id: int pad_token_id: int - bundle_uuid: str = "" @dataclass(frozen=True) diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py index 4471c5ae..685db5ad 100644 --- a/tzrec/tests/genrec_serving_contract_test.py +++ b/tzrec/tests/genrec_serving_contract_test.py @@ -136,14 +136,14 @@ def test_host_item_hash_refolds_hole_keys(self) -> None: self.assertNotEqual(_item_hash(keys), _item_hash(keys.flip(0))) def test_prompt_json_carries_the_index_builder_inputs(self) -> None: - self.assertEqual(list(self.contract), ["sid_space"]) + self.assertEqual(list(self.contract), ["sid_space", "bundle_uuid"]) space = self.contract["sid_space"] compiled = self.exported.compiled_prompt.sid_space self.assertEqual(space["base_vocab_size"], compiled.base_vocab_size) self.assertEqual(space["num_levels"], 3) self.assertEqual(space["band_lo"], list(compiled.band_lo)) self.assertEqual(space["band_hi"], list(compiled.band_hi)) - self.assertEqual(space["bundle_uuid"], "bundle-test") + self.assertEqual(self.contract["bundle_uuid"], "bundle-test") def test_tokenizer_dir_decodes_a_sid_atom(self) -> None: from transformers import AutoTokenizer diff --git a/tzrec/utils/export_util.py b/tzrec/utils/export_util.py index 7ffb06f0..fbc9400d 100644 --- a/tzrec/utils/export_util.py +++ b/tzrec/utils/export_util.py @@ -376,24 +376,6 @@ def export_model_normal( if assets is not None: for asset in assets: shutil.copy(asset, save_dir) - _export_extra_assets(model, pipeline_config, checkpoint_path, save_dir) - - -def _export_extra_assets( - model: nn.Module, - pipeline_config: EasyRecConfig, - checkpoint_path: str, - save_dir: str, -) -> None: - """Let a served module write what its runtime reads beside the scripted model. - - Discovered by duck typing through the inference wrappers, the way - checkpointing finds ``hf_backbone``; runs inside the export's own save dir - so a remote export dir is uploaded together with the scripted model. - """ - served = checkpoint_util.unwrap_to(model, "export_assets") - if served is not None: - served.export_assets(pipeline_config, checkpoint_path, save_dir) def _prepare_single_rank_distributed_embedding_export() -> bool: @@ -1902,7 +1884,6 @@ def export_distributed_embedding( merged_emb_json = _merge_sharded_embedding_json(emb_json_files) with open(os.path.join(save_dir_sparse, "sparse_embedding.json"), "w") as f: json.dump(merged_emb_json, f, indent=4) - _export_extra_assets(model, pipeline_config, checkpoint_path, save_dir) class _SparseMarkCapture(Interpreter): diff --git a/tzrec/utils/hf_export_util.py b/tzrec/utils/hf_export_util.py index 0778110d..f67c32dc 100644 --- a/tzrec/utils/hf_export_util.py +++ b/tzrec/utils/hf_export_util.py @@ -18,14 +18,25 @@ import json import os import shutil -from typing import Any, Dict, Optional, Set +import tempfile +from typing import Any, Dict, List, Optional, Set import torch from safetensors.torch import save_file from torch import nn from tzrec.constant import HF_EXPORT_META_FILENAME +from tzrec.features.feature import BaseFeature +from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.persist import ( + PROMPT_DIR, + TOKENIZER_DIR, + read_bundle_uuid, + write_serving_contract, +) +from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.utils import checkpoint_util +from tzrec.utils.filesystem_util import url_to_fs from tzrec.utils.logging_util import logger SERVING_ARCH = "PromptGenRecForCausalLM" @@ -192,3 +203,43 @@ def _derive_by_suffix( src = os.path.join(ckpt_dir, fname) if os.path.exists(src): shutil.copy(src, os.path.join(out_dir, fname)) + + +def export_hf_assets( + pipeline_config: EasyRecConfig, + features: List[BaseFeature], + checkpoint_path: str, + export_dir: str, +) -> None: + """Write what an LLM engine reads beside the scripted front-end. + + The HuggingFace weights and composite config, the extended tokenizer under + ``prompt/tokenizer`` and ``prompt/prompt.json``. A remote ``export_dir`` is + written locally and uploaded, as ``export_model`` does for its own files. + + Args: + pipeline_config: the pipeline being exported. + features: the created features the prompt compiles against. + checkpoint_path: the checkpoint the weights come from. + export_dir: the export directory. + """ + fs, local_dir = url_to_fs(export_dir) + if fs is not None: + local_dir = tempfile.mkdtemp() + dcp_to_hf(checkpoint_path, local_dir) + write_composite_config(local_dir) + prompt_config = pipeline_config.prompt_config + compiled = compile_prompt( + prompt_config, + features, + list(pipeline_config.data_config.label_fields), + tokenizer_dir=os.path.join(local_dir, PROMPT_DIR, TOKENIZER_DIR), + ) + write_serving_contract( + compiled.sid_space, + read_bundle_uuid(prompt_config.sid_space.manifest_path), + local_dir, + ) + if fs is not None: + fs.upload(local_dir, export_dir, recursive=True, file_thread_num=os.cpu_count()) + shutil.rmtree(local_dir) From 6f7ea37bd83e7ce1e51c48c9cffd60052924839b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 12:03:05 +0800 Subject: [PATCH 09/21] [ci] run the genrec pipeline as an integration test genrec_integration_test drives train_eval, eval and export through the same torchrun helpers as the rank tests and reads the export back the way the serving stack does, which retires the checkpoint-writing test fixture and the separate serving-contract test. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/main_test.py | 9 - tzrec/models/genrec_model_test.py | 38 ++- tzrec/tests/genrec_integration_test.py | 269 ++++++++++++++++++++ tzrec/tests/genrec_serving_contract_test.py | 164 ------------ tzrec/tests/prompt_integration_test.py | 213 ---------------- tzrec/tests/prompt_test_util.py | 225 +--------------- tzrec/tests/utils.py | 43 +++- 7 files changed, 350 insertions(+), 611 deletions(-) create mode 100644 tzrec/tests/genrec_integration_test.py delete mode 100644 tzrec/tests/genrec_serving_contract_test.py delete mode 100644 tzrec/tests/prompt_integration_test.py diff --git a/tzrec/main_test.py b/tzrec/main_test.py index 1359ec60..af8b497d 100644 --- a/tzrec/main_test.py +++ b/tzrec/main_test.py @@ -364,15 +364,6 @@ def test_train_loop_pairs_dense_export_with_delta_dump(self) -> None: ], ) - def test_hf_export_rejects_dense_ema(self) -> None: - # dcp_to_hf reads /model unconditionally, so it would silently - # ship raw weights where TorchScript export ships the EMA ones. - from tzrec.tests.prompt_test_util import export_tiny_genrec - - with tempfile.TemporaryDirectory() as test_dir: - with self.assertRaisesRegex(ValueError, "Dense EMA"): - export_tiny_genrec(test_dir, use_dense_ema=True) - class PredictionLifecycleTest(unittest.TestCase): """Tests for prediction lifecycle wiring.""" diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 11119842..16d01d31 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -12,6 +12,7 @@ import unittest import torch +import torch.fx from parameterized import parameterized from torchrec import KeyedJaggedTensor from transformers import AutoModelForCausalLM @@ -23,7 +24,7 @@ GenRecFrontEnd, project_slots, ) -from tzrec.models.model import ScriptWrapper +from tzrec.models.model import ScriptWrapper, TrainWrapper from tzrec.prompt.assembler import ( HOLE_POSITIONS, INPUT_IDS, @@ -39,6 +40,7 @@ _CODEBOOK, _HIST, GenRecModelTestBase, + assemble_into, create_prompt_feature, offset_sid_codes, projected_feature, @@ -212,6 +214,40 @@ def test_init_from_pretrained_replaces_the_empty_weights(self) -> None: self.assertFalse(torch.allclose(before, expected)) torch.testing.assert_close(after, expected) + def _batch_from_codes(self, hist, answer): + parsed = { + "hist.values": torch.tensor(offset_sid_codes(hist, _CODEBOOK)), + "hist.lengths": torch.tensor([len(hist)]), + "answer.values": torch.tensor(offset_sid_codes(answer, _CODEBOOK)), + "answer.lengths": torch.tensor([len(answer)]), + } + batch = Batch() + batch.additional_infos.update(assemble_into(self.compiled_prompt, parsed)) + return batch + + def test_model_resizes_to_target_vocab_size(self) -> None: + model = self._model() + rows = model.lm.get_input_embeddings().weight.shape[0] + self.assertEqual(rows, self.compiled_prompt.sid_space.target_vocab_size) + self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) + + def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: + model = self._model() + batch = self._batch_from_codes([0, 1, 2, 3, 0, 1], [1, 2, 3]) + predictions = model.predict(batch) + loss = model.loss(predictions, batch)["ce_loss"] + self.assertTrue(bool(torch.isfinite(loss))) + loss.backward() + + grad = model.lm.get_input_embeddings().weight.grad + self.assertIsNotNone(grad) + self.assertTrue(bool((grad.abs().sum() > 0))) + + def test_training_forward_survives_fx_tracing(self) -> None: + model = self._model() + + torch.fx.symbolic_trace(TrainWrapper(model)) + class GenRecFrontEndTest(GenRecModelTestBase): """The served half of the model, under the same wrapper every export uses.""" diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py new file mode 100644 index 00000000..039925d5 --- /dev/null +++ b/tzrec/tests/genrec_integration_test.py @@ -0,0 +1,269 @@ +# 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 glob +import json +import os +import shutil +import unittest +from unittest import mock + +import numpy as np +import torch +from google.protobuf import text_format +from pyarrow import parquet as pq + +from tzrec.main import _create_features +from tzrec.prompt.assembler import ( + CU_SEQLENS, + HOLE_POSITIONS, + HOLE_SLOT_COUNTS, + INPUT_IDS, + PromptAssembler, +) +from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder +from tzrec.tests import utils +from tzrec.tests.prompt_test_util import ( + _CODEBOOK, + _WORDS, + create_prompt_tokenizer, + projected_feature, +) +from tzrec.utils import config_util +from tzrec.utils.test_util import ( + create_tiny_causal_lm, + gpu_unavailable, + make_test_dir, + mark_ci_scope, +) + +_MOCK_CONFIG = "tzrec/tests/configs/genrec_causal_lm_model_mock.config" +_BUNDLE_UUID = "bundle-test" + + +class GenRecIntegrationTest(unittest.TestCase): + """train_eval -> eval -> export in subprocesses, then read the export back.""" + + def setUp(self): + self.success = False + self.test_dir = make_test_dir() + # the tiny LM trains on one rank, and one rank writes one sparse shard + patcher = mock.patch.dict(os.environ, {"TEST_NPROC_PER_NODE": "1"}) + patcher.start() + self.addCleanup(patcher.stop) + + def tearDown(self): + if self.success and os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + + def _prepare_config(self, projected: bool) -> str: + """Write the tiny backbone, tokenizer, manifest and data; return the config.""" + backbone = os.path.join(self.test_dir, "backbone") + create_tiny_causal_lm(64).save_pretrained(backbone) + tokenizer = create_prompt_tokenizer( + os.path.join(self.test_dir, "tok.json"), _WORDS + ) + manifest = os.path.join(self.test_dir, "manifest.json") + with open(manifest, "w") as f: + json.dump({"codebook": _CODEBOOK, "bundle_uuid": _BUNDLE_UUID}, f) + self.data_glob = utils.create_mock_prompt_data( + os.path.join(self.test_dir, "data"), _CODEBOOK, projected=projected + ) + + config = config_util.load_pipeline_config(_MOCK_CONFIG) + config.train_input_path = self.data_glob + config.eval_input_path = self.data_glob + config.model_config.genrec_causal_lm_model.hf_model_name_or_path = backbone + config.prompt_config.tokenizer_path = tokenizer + config.prompt_config.sid_space.manifest_path = manifest + if projected: + text_format.Merge(projected_feature("beh", 8), config.feature_configs.add()) + config.prompt_config.prompt = "History : {{hist}} . {{beh}} Predict :" + config_path = os.path.join(self.test_dir, "genrec.config") + config_util.save_message(config, config_path) + return config_path + + def _train_eval_export(self, projected: bool) -> str: + """Run the pipeline; return the trained ``pipeline.config`` path.""" + config_path = self._prepare_config(projected) + self.success = utils.test_train_eval(config_path, self.test_dir) + trained = os.path.join(self.test_dir, "pipeline.config") + if self.success: + self.success = utils.test_eval(trained, self.test_dir) + if self.success: + self.success = utils.test_export( + trained, self.test_dir, env_str="QUANT_EMB=0" + ) + self.assertTrue(self.success) + self.assertTrue( + os.path.exists(os.path.join(self.test_dir, "train/eval_result.txt")) + ) + return trained + + def _request(self, columns, rows: int = 4): + """The parsed dict a served front-end reads, from the first mock rows.""" + table = pq.read_table(sorted(glob.glob(self.data_glob))[0]).slice(0, rows) + out = {} + for column in columns: + lists = table.column(column).to_pylist() + out[column + ".values"] = torch.tensor( + [v for row in lists for v in row], dtype=torch.int64 + ) + out[column + ".lengths"] = torch.tensor([len(row) for row in lists]) + return out + + def test_genrec_train_eval_export(self): + trained = self._train_eval_export(projected=True) + export_dir = os.path.join(self.test_dir, "export") + for name in ( + "scripted_model.pt", + "fg.json", + "pipeline.config", + "model_acc.json", + "config.json", + "model.safetensors", + "prompt/prompt.json", + "prompt/tokenizer/tokenizer.json", + "prompt/tokenizer/tokenizer_config.json", + ): + self.assertTrue(os.path.exists(os.path.join(export_dir, name)), name) + self.assertFalse(os.path.exists(os.path.join(export_dir, "sparse"))) + + # the artifact is the collator's walk plus the serving-only fold + config = config_util.load_pipeline_config(trained) + features = _create_features(list(config.feature_configs), config.data_config) + compiled = compile_prompt(config.prompt_config, features, ["answer"]) + data = self._request(["hist", "beh"]) + front_end = torch.jit.load(os.path.join(export_dir, "scripted_model.pt")) + # an LLM engine calls it with the batch alone; a processor adds a device + out = front_end(data) + with_device = front_end(data, torch.device("cpu")) + walk = PromptAssembler( + compiled.prompt_plan, compiled.sid_space, include_response=False + )(data) + for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, HOLE_SLOT_COUNTS): + self.assertTrue(torch.equal(out[key], walk[key]), key) + self.assertTrue(torch.equal(with_device[key], walk[key]), key) + self.assertTrue( + torch.equal(out[HOLE_KEYS], HoleKeyBuilder(compiled.prompt_plan)(data)) + ) + positions = out[HOLE_POSITIONS] + self.assertGreater(positions.numel(), 0) + self.assertTrue( + bool( + torch.all( + out[INPUT_IDS][positions] == compiled.sid_space.sentinel_token_id + ) + ) + ) + self.assertEqual(tuple(out["slot_embeds"].shape), (int(positions.numel()), 32)) + self.assertEqual(out["slot_embeds"].dtype, torch.float32) + + # what the engine loads beside it + with open(os.path.join(export_dir, "config.json"), "r") as f: + hf_config = json.load(f) + self.assertEqual(hf_config["architectures"], ["PromptGenRecForCausalLM"]) + self.assertEqual(hf_config["model_type"], "prompt_genrec") + self.assertEqual(hf_config["text_config"]["model_type"], "qwen2") + self.assertEqual( + hf_config["text_config"]["vocab_size"], compiled.sid_space.target_vocab_size + ) + with open(os.path.join(export_dir, "prompt", "prompt.json"), "r") as f: + contract = json.load(f) + self.assertEqual(list(contract), ["sid_space", "bundle_uuid"]) + self.assertEqual(contract["bundle_uuid"], _BUNDLE_UUID) + self.assertEqual( + contract["sid_space"]["band_lo"], list(compiled.sid_space.band_lo) + ) + self.assertEqual( + contract["sid_space"]["band_hi"], list(compiled.sid_space.band_hi) + ) + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained( + os.path.join(export_dir, "prompt", "tokenizer") + ) + self.assertEqual(tokenizer.decode([compiled.sid_space.band_lo[0]]), "<|sid_0|>") + self.assertEqual( + tokenizer.convert_ids_to_tokens(compiled.sid_space.sentinel_token_id), + "<|pg_hole|>", + ) + self.assertEqual(tokenizer.eos_token_id, compiled.sid_space.eos_token_id) + + # Dense EMA weights never reach the HF conversion + config.export_config.use_dense_ema = True + ema_config = os.path.join(self.test_dir, "dense_ema.config") + config_util.save_message(config, ema_config) + self.assertFalse( + utils.test_export( + ema_config, self.test_dir, export_dir=os.path.join(self.test_dir, "ema") + ) + ) + with open(os.path.join(self.test_dir, "log_export.txt"), "r") as f: + self.assertIn("Dense EMA", f.read()) + + @unittest.skipIf(*gpu_unavailable) + @mark_ci_scope("gpu") + def test_genrec_export_distributed_embedding(self): + trained = self._train_eval_export(projected=True) + dist_dir = os.path.join(self.test_dir, "export_dist") + self.success = utils.test_export( + trained, + self.test_dir, + export_dir=dist_dir, + env_str="USE_DISTRIBUTED_EMBEDDING=1 QUANT_EMB=0", + ) + self.assertTrue(self.success) + sparse_dir = os.path.join(dist_dir, "sparse") + for name in ( + "sparse_embeddings-00-of-01.npz", + "sparse_embedding.json", + "sparse_features.json", + ): + self.assertTrue(os.path.exists(os.path.join(sparse_dir, name)), name) + for name in ("config.json", "model.safetensors", "prompt/prompt.json"): + self.assertTrue(os.path.exists(os.path.join(dist_dir, name)), name) + with open(os.path.join(dist_dir, "model_acc.json"), "r") as f: + self.assertEqual(json.load(f)["DISTRIBUTED_EMBEDDING"], "1") + with open(os.path.join(dist_dir, "dense_meta.json"), "r") as f: + self.assertEqual(json.load(f)["sequence__ec"], ["beh__ec", "beh__lengths"]) + + # a processor simulator: one request is one user, looked up in the + # exported tables the way the distributed-embedding stage does, then + # fed to the dense stage input-tiled with one candidate + request = self._request(["hist", "beh"], rows=1) + with open(os.path.join(sparse_dir, "sparse_features.json"), "r") as f: + table_name = json.load(f)["beh__ec"]["embedding_name"] + with np.load(os.path.join(sparse_dir, "sparse_embeddings-00-of-01.npz")) as npz: + table = npz[table_name] + data = dict(request) + data["beh"] = torch.from_numpy( + table[request["beh.values"].numpy()].astype(np.float32) + ) + data["beh__lengths"] = request["beh.lengths"] + data["batch_size"] = torch.tensor(1) + # the planner exported on the GPU; the processor loads onto its device + got = torch.jit.load( + os.path.join(dist_dir, "scripted_model.pt"), map_location="cpu" + )(data) + expected = torch.jit.load( + os.path.join(self.test_dir, "export", "scripted_model.pt") + )(request) + self.assertTrue( + torch.allclose(got["slot_embeds"], expected["slot_embeds"], atol=1e-6) + ) + self.assertTrue(torch.equal(got["hole_keys"], expected["hole_keys"])) + self.assertEqual(got["input_ids"].tolist(), expected["input_ids"].tolist()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/tests/genrec_serving_contract_test.py b/tzrec/tests/genrec_serving_contract_test.py deleted file mode 100644 index 685db5ad..00000000 --- a/tzrec/tests/genrec_serving_contract_test.py +++ /dev/null @@ -1,164 +0,0 @@ -# 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. - -"""Reads a genrec export the way the SGLang genrec stack does. - -No sglang import: these tests pin the contract from the consumer's side -- -the composite config its model wrapper resolves, the front-end outputs its -multimodal processor cuts into items, the host-side item hash it re-folds -from ``hole_keys``, and the ``prompt.json`` fields its constraint-index -builder reads -- so a change on this side that would break serving fails here. -""" - -import io -import json -import os -import unittest - -import numpy as np -import torch -from safetensors.torch import load_file - -from tzrec.prompt.hole_keys import mix64 -from tzrec.tests.prompt_test_util import export_tiny_genrec -from tzrec.utils.test_util import make_test_dir - -# sglang's multimodal/processors/prompt_genrec.py folds a slot's per-hole keys -# into one item hash with this mixer and this position constant -_C_POSITION = -6752110988234923001 -_MASK64 = (1 << 64) - 1 - - -def _mix64_host(value: int) -> int: - z = (value * 0x9E3779B97F4A7C15) & _MASK64 - z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & _MASK64 - z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & _MASK64 - return z ^ (z >> 31) - - -def _item_hash(keys) -> int: - total = 0 - for index, key in enumerate(keys.tolist()): - total = (total + _mix64_host((key + _C_POSITION * index) & _MASK64)) & _MASK64 - return total - - -class GenRecServingContractTest(unittest.TestCase): - @classmethod - def setUpClass(cls) -> None: - cls.exported = export_tiny_genrec(make_test_dir(), bundle_uuid="bundle-test") - with open(os.path.join(cls.exported.export_dir, "config.json"), "r") as f: - cls.config = json.load(f) - with open( - os.path.join(cls.exported.export_dir, "prompt", "prompt.json"), "r" - ) as f: - cls.contract = json.load(f) - cls.front_end = torch.jit.load( - os.path.join(cls.exported.export_dir, "scripted_model.pt") - ) - - def _payload(self): - """A request as the in-process processor receives it: an npz blob.""" - buffer = io.BytesIO() - np.savez_compressed( - buffer, **{k: v.numpy() for k, v in self.exported.sample_data.items()} - ) - buffer.seek(0) - with np.load(buffer, allow_pickle=False) as data: - return {key: torch.from_numpy(np.asarray(data[key])) for key in data.files} - - def test_config_is_composite_and_names_the_backbone(self) -> None: - self.assertEqual(self.config["architectures"], ["PromptGenRecForCausalLM"]) - self.assertEqual(self.config["model_type"], "prompt_genrec") - text_config = self.config["text_config"] - self.assertEqual(text_config["model_type"], "qwen2") - self.assertEqual( - text_config["vocab_size"], - self.exported.compiled_prompt.sid_space.target_vocab_size, - ) - - def test_weights_keep_the_backbones_own_names(self) -> None: - """The wrapper delegates load_weights wholesale, so no remapping exists.""" - from transformers import AutoConfig, AutoModelForCausalLM - - config = AutoConfig.for_model(**self.config["text_config"]) - with torch.device("meta"): - backbone = AutoModelForCausalLM.from_config(config) - expected = set(backbone.state_dict().keys()) - exported = set( - load_file(os.path.join(self.exported.export_dir, "model.safetensors")) - ) - self.assertTrue(exported <= expected, exported - expected) - self.assertIn("model.embed_tokens.weight", exported) - - def test_front_end_output_feeds_the_processor(self) -> None: - # sglang calls it with the batch alone; the processor adds a device - out = self.front_end(self._payload()) - with_device = self.front_end(self._payload(), torch.device("cpu")) - self.assertTrue(torch.equal(with_device["hole_keys"], out["hole_keys"])) - for key in ( - "input_ids", - "hole_positions", - "slot_embeds", - "hole_keys", - "hole_slot_counts", - ): - self.assertIn(key, out) - positions = out["hole_positions"] - self.assertEqual(positions.dtype, torch.int64) - self.assertTrue(bool(torch.all(positions[1:] > positions[:-1]))) - self.assertEqual(int(out["hole_slot_counts"].sum()), int(positions.numel())) - self.assertEqual( - tuple(out["slot_embeds"].shape), - (int(positions.numel()), self.config["text_config"]["hidden_size"]), - ) - self.assertEqual(out["slot_embeds"].dtype, torch.float32) - sentinel = self.contract["sid_space"]["sentinel_token_id"] - self.assertTrue(bool(torch.all(out["input_ids"][positions] == sentinel))) - self.assertEqual(out["hole_keys"].dtype, torch.int64) - - def test_host_item_hash_refolds_hole_keys(self) -> None: - """The host's per-item outer fold reuses the front-end's mixer.""" - for value in (0, 1, -1, 123456789, -(2**40)): - signed = _mix64_host(value & _MASK64) - signed = signed - (1 << 64) if signed >= 1 << 63 else signed - self.assertEqual(int(mix64(torch.tensor([value]))[0]), signed) - keys = self.front_end(self._payload())["hole_keys"] - self.assertEqual(_item_hash(keys), _item_hash(keys)) - self.assertNotEqual(_item_hash(keys), _item_hash(keys.flip(0))) - - def test_prompt_json_carries_the_index_builder_inputs(self) -> None: - self.assertEqual(list(self.contract), ["sid_space", "bundle_uuid"]) - space = self.contract["sid_space"] - compiled = self.exported.compiled_prompt.sid_space - self.assertEqual(space["base_vocab_size"], compiled.base_vocab_size) - self.assertEqual(space["num_levels"], 3) - self.assertEqual(space["band_lo"], list(compiled.band_lo)) - self.assertEqual(space["band_hi"], list(compiled.band_hi)) - self.assertEqual(self.contract["bundle_uuid"], "bundle-test") - - def test_tokenizer_dir_decodes_a_sid_atom(self) -> None: - from transformers import AutoTokenizer - - tokenizer = AutoTokenizer.from_pretrained( - os.path.join(self.exported.export_dir, "prompt", "tokenizer") - ) - space = self.contract["sid_space"] - self.assertEqual(tokenizer.decode([space["band_lo"][0]]), "<|sid_0|>") - self.assertEqual( - tokenizer.convert_ids_to_tokens(space["sentinel_token_id"]), "<|pg_hole|>" - ) - self.assertEqual(tokenizer.eos_token_id, space["eos_token_id"]) - self.assertEqual(tokenizer.pad_token_id, space["pad_token_id"]) - - -if __name__ == "__main__": - unittest.main() diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py deleted file mode 100644 index 4f63bad6..00000000 --- a/tzrec/tests/prompt_integration_test.py +++ /dev/null @@ -1,213 +0,0 @@ -# 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 json -import os -import subprocess -import sys -import unittest - -import numpy as np -import torch -import torch.fx - -from tzrec.datasets.utils import Batch -from tzrec.models.model import TrainWrapper -from tzrec.prompt.assembler import ( - CU_SEQLENS, - HOLE_POSITIONS, - INPUT_IDS, - PromptAssembler, -) -from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder -from tzrec.tests.prompt_test_util import ( - GenRecModelTestBase, - assemble_into, - export_tiny_genrec, - offset_sid_codes, -) -from tzrec.utils.test_util import gpu_unavailable, make_test_dir, mark_ci_scope - -_CODEBOOK = [4, 4, 4] -_WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] - - -class PromptStackIntegrationTest(GenRecModelTestBase): - """compile -> assemble -> model, on the real code path.""" - - def _batch_from_codes(self, hist, answer): - parsed = { - "hist.values": torch.tensor(offset_sid_codes(hist, _CODEBOOK)), - "hist.lengths": torch.tensor([len(hist)]), - "answer.values": torch.tensor(offset_sid_codes(answer, _CODEBOOK)), - "answer.lengths": torch.tensor([len(answer)]), - } - batch = Batch() - batch.additional_infos.update(assemble_into(self.compiled_prompt, parsed)) - return batch - - def test_model_resizes_to_target_vocab_size(self) -> None: - model = self._model() - rows = model.lm.get_input_embeddings().weight.shape[0] - self.assertEqual(rows, self.compiled_prompt.sid_space.target_vocab_size) - self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) - - def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: - model = self._model() - batch = self._batch_from_codes([0, 1, 2, 3, 0, 1], [1, 2, 3]) - predictions = model.predict(batch) - loss = model.loss(predictions, batch)["ce_loss"] - self.assertTrue(bool(torch.isfinite(loss))) - loss.backward() - - grad = model.lm.get_input_embeddings().weight.grad - self.assertIsNotNone(grad) - self.assertTrue(bool((grad.abs().sum() > 0))) - - def test_training_forward_survives_fx_tracing(self) -> None: - model = self._model() - - torch.fx.symbolic_trace(TrainWrapper(model)) - - -class GenRecExportIntegrationTest(unittest.TestCase): - """checkpoint -> export -> the artifacts an LLM engine and a processor load.""" - - def setUp(self) -> None: - self.test_dir = make_test_dir() - - def test_export_writes_one_serving_directory(self) -> None: - exported = export_tiny_genrec(self.test_dir, env={"QUANT_EMB": "0"}) - for name in ( - "scripted_model.pt", - "fg.json", - "pipeline.config", - "model_acc.json", - "config.json", - "model.safetensors", - "prompt/prompt.json", - "prompt/tokenizer/tokenizer.json", - "prompt/tokenizer/tokenizer_config.json", - ): - self.assertTrue( - os.path.exists(os.path.join(exported.export_dir, name)), name - ) - self.assertFalse(os.path.exists(os.path.join(exported.export_dir, "sparse"))) - - # the exported artifact and the collator are two call sites of one walk; - # only the artifact folds hole keys - data = exported.sample_data - compiled = exported.compiled_prompt - collator = PromptAssembler( - compiled.prompt_plan, compiled.sid_space, include_response=False - )(data) - front_end = torch.jit.load( - os.path.join(exported.export_dir, "scripted_model.pt") - ) - out = front_end(data, torch.device("cpu")) - for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS): - self.assertTrue(torch.equal(out[key], collator[key]), key) - self.assertTrue( - torch.equal(out[HOLE_KEYS], HoleKeyBuilder(compiled.prompt_plan)(data)) - ) - self.assertNotIn(HOLE_KEYS, collator) - self.assertEqual( - tuple(out["slot_embeds"].shape), (int(out[HOLE_POSITIONS].numel()), 32) - ) - - @unittest.skipIf(*gpu_unavailable) - @mark_ci_scope("gpu") - def test_distributed_embedding_export_writes_the_processor_shape(self) -> None: - exported = export_tiny_genrec(self.test_dir, env={"QUANT_EMB": "0"}) - dist_dir = os.path.join(self.test_dir, "export_dist") - # its own process: the planner wants a fresh nccl group on the device - env = dict(os.environ) - env.update( - { - "PYTHONPATH": ".", - "USE_DISTRIBUTED_EMBEDDING": "1", - "QUANT_EMB": "0", - "MASTER_ADDR": "127.0.0.1", - "MASTER_PORT": os.environ.get("MASTER_PORT", "29511"), - "RANK": "0", - "LOCAL_RANK": "0", - "WORLD_SIZE": "1", - } - ) - log_path = os.path.join(self.test_dir, "export_dist.log") - with open(log_path, "w") as log: - code = subprocess.call( - [ - sys.executable, - "-m", - "tzrec.export", - "--pipeline_config_path", - exported.config_path, - "--export_dir", - dist_dir, - "--checkpoint_path", - exported.checkpoint_dir, - ], - env=env, - stdout=log, - stderr=subprocess.STDOUT, - ) - with open(log_path, "r") as log: - self.assertEqual(code, 0, log.read()[-4000:]) - sparse_dir = os.path.join(dist_dir, "sparse") - for name in ( - "sparse_embeddings-00-of-01.npz", - "sparse_embedding.json", - "sparse_features.json", - ): - self.assertTrue(os.path.exists(os.path.join(sparse_dir, name)), name) - with open(os.path.join(dist_dir, "model_acc.json"), "r") as f: - self.assertEqual(json.load(f)["DISTRIBUTED_EMBEDDING"], "1") - with open(os.path.join(dist_dir, "dense_meta.json"), "r") as f: - dense_meta = json.load(f) - self.assertEqual(dense_meta["sequence__ec"], ["beh__ec", "beh__lengths"]) - - # a processor simulator: one request is one user, looked up in the - # exported tables the way the distributed-embedding stage does, then - # fed to the dense stage input-tiled with one candidate - sample = exported.sample_data - request = { - "hist.values": sample["hist.values"][: int(sample["hist.lengths"][0])], - "hist.lengths": sample["hist.lengths"][:1], - "beh.values": sample["beh.values"][: int(sample["beh.lengths"][0])], - "beh.lengths": sample["beh.lengths"][:1], - } - with open(os.path.join(sparse_dir, "sparse_features.json"), "r") as f: - table_name = json.load(f)["beh__ec"]["embedding_name"] - with np.load(os.path.join(sparse_dir, "sparse_embeddings-00-of-01.npz")) as npz: - table = npz[table_name] - data = dict(request) - data["beh"] = torch.from_numpy( - table[request["beh.values"].numpy()].astype(np.float32) - ) - data["beh__lengths"] = request["beh.lengths"] - data["batch_size"] = torch.tensor(1) - # the planner exported on the GPU; the processor loads onto its device - got = torch.jit.load( - os.path.join(dist_dir, "scripted_model.pt"), map_location="cpu" - )(data) - expected = torch.jit.load( - os.path.join(exported.export_dir, "scripted_model.pt") - )(request) - self.assertTrue( - torch.allclose(got["slot_embeds"], expected["slot_embeds"], atol=1e-6) - ) - self.assertTrue(torch.equal(got["hole_keys"], expected["hole_keys"])) - self.assertEqual(got["input_ids"].tolist(), expected["input_ids"].tolist()) - - -if __name__ == "__main__": - unittest.main() diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index cfc7635b..551f34d5 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -9,31 +9,24 @@ # See the License for the specific language governing permissions and # limitations under the License. -import dataclasses -import json import os import unittest -from typing import Any, Dict, List, Optional, Sequence +from typing import Any, Dict, Sequence import numpy as np import torch from google.protobuf import text_format from tokenizers import Tokenizer, models, pre_tokenizers -from torch import distributed as dist from tzrec.datasets.utils import BASE_DATA_GROUP, Batch from tzrec.features.feature import BaseFeature, FgMode, create_features -from tzrec.main import _create_features, _create_model, export -from tzrec.models.model import TrainWrapper +from tzrec.main import _create_model from tzrec.prompt.assembler import PROMPT_INFO_PREFIX, PromptAssembler from tzrec.prompt.compile import compile_prompt from tzrec.prompt.types import CompiledPrompt from tzrec.protos import feature_pb2 from tzrec.protos.model_pb2 import ModelConfig -from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.protos.prompt_pb2 import PromptConfig -from tzrec.utils.hf_export_util import write_hf_assets -from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import create_tiny_causal_lm, make_test_dir @@ -169,217 +162,3 @@ def _batch(self, parsed, compiled_prompt=None, sparse=None): batch = Batch(sparse_features={BASE_DATA_GROUP: sparse} if sparse else {}) batch.additional_infos.update(streams) return batch - - -def write_genrec_checkpoint(model: torch.nn.Module, ckpt_dir: str) -> str: - """Save a model the way a training run does, minus the dynamic-table dump. - - Args: - model: the bare genrec model; wrapped and initialized here. - ckpt_dir: the ``model.ckpt-N`` directory to write. - - Returns: - ``ckpt_dir``. - """ - from torch.distributed.checkpoint import save - - wrapped = TrainWrapper(model) - init_parameters(wrapped, torch.device("cpu")) - save(wrapped.state_dict(), checkpoint_id=os.path.join(ckpt_dir, "model")) - write_hf_assets(wrapped, ckpt_dir) - return ckpt_dir - - -@dataclasses.dataclass -class ExportedGenRec: - """A tiny genrec model, its checkpoint and its export. - - Args: - config: the pipeline config the export ran on. - config_path: where it was written. - features: the created features. - compiled_prompt: the prompt as the export compiled it. - checkpoint_dir: the checkpoint the export converted. - export_dir: the export directory. - sample_data: one parsed batch as the served module reads it, the dict - ``Batch.to_dict`` emits for the mock data. - """ - - config: EasyRecConfig - config_path: str - features: List[BaseFeature] - compiled_prompt: CompiledPrompt - checkpoint_dir: str - export_dir: str - sample_data: Dict[str, torch.Tensor] - - -def _write_mock_prompt_data(path: str, projected: bool, rows: int = 8) -> str: - """Write a parquet of SID histories the mock prompt reads. - - Returns: - The glob a data config points at. - """ - import pyarrow as pa - import pyarrow.parquet as pq - - rng = np.random.default_rng(0) - columns = { - "hist": [ - offset_sid_codes(rng.integers(0, 4, size=6), _CODEBOOK).tolist() - for _ in range(rows) - ], - "answer": [ - offset_sid_codes(rng.integers(0, 4, size=3), _CODEBOOK).tolist() - for _ in range(rows) - ], - } - if projected: - columns["beh"] = [rng.integers(0, 32, size=2).tolist() for _ in range(rows)] - os.makedirs(path, exist_ok=True) - pq.write_table( - pa.table( - {k: pa.array(v, type=pa.list_(pa.int64())) for k, v in columns.items()} - ), - os.path.join(path, "part-0.parquet"), - ) - return os.path.join(path, "*.parquet") - - -def export_tiny_genrec( - test_dir: str, - projected: bool = True, - bundle_uuid: Optional[str] = None, - use_dense_ema: bool = False, - env: Optional[Dict[str, str]] = None, -) -> ExportedGenRec: - """Train nothing, save a checkpoint of a tiny genrec model, and export it. - - The prompt carries an INLINE SID history and, when ``projected``, one - PROJECTED behaviour slot, which is what makes the export carry a table and - a projection. The export traces on one batch of mock parquet, as any tzrec - export does. - - Args: - test_dir: scratch directory. - projected: whether to add the projected slot. - bundle_uuid: when given, a SID manifest carrying it is written and - referenced, so the export records a bundle identity. - use_dense_ema: ask the export for Dense EMA weights, which the HF - export refuses. - env: extra environment for the export, e.g. ``USE_DISTRIBUTED_EMBEDDING``. - - Returns: - Everything a test needs to read the export back. - """ - from unittest import mock - - from tzrec.constant import Mode - from tzrec.datasets.dataset import create_dataloader - - backbone = os.path.join(test_dir, "backbone") - create_tiny_causal_lm(64).save_pretrained(backbone) - tok = create_prompt_tokenizer(os.path.join(test_dir, "tok.json"), _WORDS) - manifest = "" - if bundle_uuid is not None: - manifest = os.path.join(test_dir, "manifest.json") - with open(manifest, "w") as f: - json.dump({"codebook": _CODEBOOK, "bundle_uuid": bundle_uuid}, f) - data_glob = _write_mock_prompt_data(os.path.join(test_dir, "data"), projected) - - config = EasyRecConfig() - text_format.Merge( - f''' -train_input_path: "{data_glob}" eval_input_path: "{data_glob}" -model_dir: "{test_dir}/train" -train_config {{ - sparse_optimizer {{ adagrad_optimizer {{ lr: 0.0 }} constant_learning_rate {{}} }} - dense_optimizer {{ adam_optimizer {{ lr: 0.0001 }} constant_learning_rate {{}} }} - num_epochs: 1 -}} -data_config {{ - batch_size: 4 dataset_type: ParquetDataset fg_mode: FG_NONE - label_fields: "answer" num_workers: 1 -}} -{"export_config { use_dense_ema: true }" if use_dense_ema else ""} -feature_configs {{ {_HIST} }} -{"feature_configs { " + projected_feature("beh", 8) + " }" if projected else ""} -prompt_config {{ - tokenizer_path: "{tok}" - prompt: "History : {{{{hist}}}} .{" {{beh}}" if projected else ""} Predict :" - response: "{{{{answer}}}}" - sid_space {{ - codebook: 4 codebook: 4 codebook: 4 - {f'manifest_path: "{manifest}"' if manifest else ""} - }} - max_length: 64 -}} -model_config {{ - genrec_causal_lm_model {{ - hf_model_name_or_path: "{backbone}" - common {{ beam_widths: 2 beam_widths: 2 beam_widths: 2 num_return_sequences: 2 }} - }} -}} -''', - config, - ) - os.makedirs(config.model_dir) - config_path = os.path.join(test_dir, "pipeline.config") - with open(config_path, "w") as f: - f.write(text_format.MessageToString(config)) - - features = _create_features(list(config.feature_configs), config.data_config) - compiled_prompt = compile_prompt(config.prompt_config, features, ["answer"]) - model = _create_model( - config.model_config, features, ["answer"], compiled_prompt=compiled_prompt - ) - checkpoint_dir = write_genrec_checkpoint( - model, os.path.join(config.model_dir, "model.ckpt-1") - ) - export_dir = os.path.join(test_dir, "export") - # export_model_normal opens a single-process gloo group; close it again so - # a later test in this process can open its own - group_was_open = dist.is_initialized() - with mock.patch.dict( - os.environ, - { - "MASTER_ADDR": "127.0.0.1", - "MASTER_PORT": os.environ.get("MASTER_PORT", str(_free_port())), - "RANK": "0", - "LOCAL_RANK": "0", - "WORLD_SIZE": "1", - **(env or {}), - }, - ): - try: - export(config_path, export_dir, checkpoint_path=checkpoint_dir) - finally: - if dist.is_initialized() and not group_was_open: - dist.destroy_process_group() - dataloader = create_dataloader( - config.data_config, features, data_glob, mode=Mode.PREDICT - ) - batch = next(iter(dataloader)) - sample_data = { - k: v - for k, v in batch.to_dict(sparse_dtype=torch.int64).items() - if not k.startswith(PROMPT_INFO_PREFIX) - } - return ExportedGenRec( - config=config, - config_path=config_path, - features=features, - compiled_prompt=compiled_prompt, - checkpoint_dir=checkpoint_dir, - export_dir=export_dir, - sample_data=sample_data, - ) - - -def _free_port() -> int: - """A port the single-process gloo group can bind.""" - import socket - - with socket.socket() as sock: - sock.bind(("127.0.0.1", 0)) - return int(sock.getsockname()[1]) diff --git a/tzrec/tests/utils.py b/tzrec/tests/utils.py index 7a760038..06c97a1e 100644 --- a/tzrec/tests/utils.py +++ b/tzrec/tests/utils.py @@ -14,12 +14,13 @@ import os import random from collections import OrderedDict, defaultdict -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Sequence, Tuple import numpy as np import numpy.typing as npt import pyarrow as pa import pyarrow.dataset as ds +import pyarrow.parquet as pq import torch from tzrec.acc.utils import is_aot_predict, is_trt_predict @@ -533,6 +534,46 @@ def create_mock_data( return os.path.join(data_dir, f"*.{fmt}"), t +def create_mock_prompt_data( + path: str, codebook: Sequence[int], projected: bool, num_rows: int = 8 +) -> str: + """Write a parquet of SID histories for the prompt-native tests. + + ``hist`` and ``answer`` carry offset SID codes, ``level_offsets[l] + code``, + which is what a prompt reads; ``beh`` is an id sequence for a PROJECTED slot. + + Args: + path: directory to write into. + codebook: per-level SID vocabulary sizes. + projected: whether to add the ``beh`` column. + num_rows: samples to write. + + Returns: + The glob a data config points at. + """ + rng = np.random.default_rng(0) + offsets = np.cumsum([0, *codebook[:-1]]) + + def codes(items: int) -> List[int]: + drawn = rng.integers(0, codebook, size=(items, len(codebook))) + return (drawn + offsets).reshape(-1).tolist() + + columns = { + "hist": [codes(2) for _ in range(num_rows)], + "answer": [codes(1) for _ in range(num_rows)], + } + if projected: + columns["beh"] = [rng.integers(0, 32, size=2).tolist() for _ in range(num_rows)] + os.makedirs(path, exist_ok=True) + pq.write_table( + pa.table( + {k: pa.array(v, type=pa.list_(pa.int64())) for k, v in columns.items()} + ), + os.path.join(path, "part-0.parquet"), + ) + return os.path.join(path, "*.parquet") + + def create_mock_join_data( data_dir: str, inputs: Dict[str, MockInput], From 5e4516896c0b09cf7bd904a6848e0ca899a4624f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 13:38:31 +0800 Subject: [PATCH 10/21] [ci] mark the GPU-only hole-key fold test with its CI scope Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/prompt/hole_keys_test.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tzrec/prompt/hole_keys_test.py b/tzrec/prompt/hole_keys_test.py index 6e4752ac..22bb3cdc 100644 --- a/tzrec/prompt/hole_keys_test.py +++ b/tzrec/prompt/hole_keys_test.py @@ -28,6 +28,7 @@ ) from tzrec.protos.model_pb2 import FeatureGroupType from tzrec.utils.fx_util import symbolic_trace +from tzrec.utils.test_util import gpu_unavailable, mark_ci_scope def _slot( @@ -243,7 +244,8 @@ def forward(self, batch): ) self.assertTrue(torch.equal(traced(batch), self.module(batch))) - @unittest.skipIf(not torch.cuda.is_available(), "no GPU") + @unittest.skipIf(*gpu_unavailable) + @mark_ci_scope("gpu") def test_the_fold_is_bit_identical_across_devices(self) -> None: """Integer addition cannot depend on the order a device reduces in.""" batch = _tensors( From c5fc3af7b504acf139f6adda80a2aa245cd68782 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 14:05:04 +0800 Subject: [PATCH 11/21] [doc] shorten the genrec module docstrings Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/models/genrec_model.py | 13 +++++-------- tzrec/prompt/assembler.py | 16 ++++++---------- tzrec/prompt/hole_keys.py | 25 +++++++------------------ tzrec/tests/genrec_integration_test.py | 12 ------------ 4 files changed, 18 insertions(+), 48 deletions(-) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 8d57ab99..6fe3e9fb 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -11,14 +11,11 @@ """Shared causal-LM plumbing for generative recommendation models. -This layer builds an empty causal LM, resizes its vocabulary, wires slot -projections, converts SID coordinate systems and scores the response window. -A family subclass owns its forward and decode path. - -``GenRecFrontEnd`` is the half of the model tzrec serves: the assembled prompt -and the projected slots, everything before the LM's embedding gather. It is -exported like any tzrec model; the LM itself is handed to an LLM engine as the -HuggingFace weights beside it. +Builds the causal LM, resizes its vocabulary, wires slot projections, converts +SID coordinates and scores the response window; a family subclass owns its +forward and decode path. ``GenRecFrontEnd`` is the served half, the assembled +prompt and the projected slots, exported like any tzrec model beside the LM's +HuggingFace weights. """ import inspect diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index fdae9824..eb0d58cc 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -11,16 +11,12 @@ """The prompt assembler: one scripted walk with two call sites. -It runs in the dataloader worker after the features are parsed, and again -inside the exported front-end at serving. ``PromptPlan`` is a compile-time -constant, so the segment loop unrolls into parallel constant lists at -construction and what remains is jagged integer arithmetic -- ``cumsum``, -``repeat_interleave``, ``index_copy_`` -- with no data-dependent control flow, -which is what lets ``torch.jit.script`` carry the same module into a runtime -that has no tzrec source. Serving therefore never reimplements this walk. - -The walk is prompt structure only. The prefix-cache identity of each hole, -``hole_keys``, is a serving concern that ``hole_keys.py`` folds beside it. +The collator runs it after parsing and the exported front-end runs the same +module at serving. ``PromptPlan`` unrolls into constant lists at construction, +leaving jagged integer arithmetic with no data-dependent control flow, which +is what ``torch.jit.script`` can carry into a runtime without tzrec source. +The walk is prompt structure only; ``hole_keys.py`` folds the prefix-cache +identity of each hole beside it. """ from typing import Dict, Final, List, Tuple diff --git a/tzrec/prompt/hole_keys.py b/tzrec/prompt/hole_keys.py index 4f3a3f13..05830f32 100644 --- a/tzrec/prompt/hole_keys.py +++ b/tzrec/prompt/hole_keys.py @@ -11,24 +11,13 @@ """The ``hole_keys`` fold: a serving-only prefix-cache identity per hole. -An LLM engine keys its prefix cache on token ids, and a PROJECTED slot writes -the same sentinel at the same position for every request, so the engine needs -a per-hole identity to match on instead. ``HoleKeyBuilder`` folds each hole's -*input* values -- never the projected vector -- into one ``int64`` key per -hole, row-aligned with the assembler's ``hole_positions``: the projected -occurrences in emission order, then the samples. Only the serving wrapper -composes it; training never computes a key it would discard. - -The fold is integer end to end: ``int64`` addition is associative and -commutative and wraps deterministically, so a key cannot depend on the order -``index_add_`` happens to reduce in, on the device, or on how the batch was -split. A float accumulator would satisfy "do not fold the projected vector" in -letter and reintroduce the variance in spirit. - -No plan or model identity enters a key. A cache shared across model versions -already matches their static tokens alike, so only the engine can namespace -it, and the engine folds the keys it receives with a position constant of its -own. Nothing outside the scripted front-end reads the constants below. +A PROJECTED slot writes the same sentinel at the same position for every +request, so an engine keying its prefix cache on token ids needs a per-hole +identity instead. ``HoleKeyBuilder`` folds each hole's input values, never the +projected vector, into one ``int64`` key per hole, row-aligned with the +assembler's ``hole_positions``. The fold is integer end to end, so a key cannot +depend on reduction order, device or batch split. No plan or model identity +enters a key: only the engine can namespace a shared cache. """ from typing import Dict, Final, List diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 039925d5..5c103d43 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -199,18 +199,6 @@ def test_genrec_train_eval_export(self): ) self.assertEqual(tokenizer.eos_token_id, compiled.sid_space.eos_token_id) - # Dense EMA weights never reach the HF conversion - config.export_config.use_dense_ema = True - ema_config = os.path.join(self.test_dir, "dense_ema.config") - config_util.save_message(config, ema_config) - self.assertFalse( - utils.test_export( - ema_config, self.test_dir, export_dir=os.path.join(self.test_dir, "ema") - ) - ) - with open(os.path.join(self.test_dir, "log_export.txt"), "r") as f: - self.assertIn("Dense EMA", f.read()) - @unittest.skipIf(*gpu_unavailable) @mark_ci_scope("gpu") def test_genrec_export_distributed_embedding(self): From 03a4e565e90a87214e8b3e0ffaf6ab435f36b24d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 14:44:33 +0800 Subject: [PATCH 12/21] [refactor] trust the parsed batch in the prompt walk The walk and the fold read the parsed dict the way every scripted module does: no schema, width, band, length or batch-size checks, each of which cost a device sync at serving, and no prompt_ prefix on the batch keys, since nothing else in additional_infos shares those names. The batch size comes from the first slot member, and a dense sequence member now folds per row instead of being refused. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/datasets/dataset.py | 6 +- tzrec/models/genrec_causal_lm_model.py | 16 +- tzrec/models/genrec_causal_lm_model_test.py | 16 +- tzrec/models/genrec_model.py | 22 +- tzrec/models/genrec_model_test.py | 6 +- tzrec/models/model.py | 10 +- tzrec/prompt/assembler.py | 219 +++----------------- tzrec/prompt/assembler_test.py | 137 +++--------- tzrec/prompt/hole_keys.py | 107 ++++------ tzrec/prompt/hole_keys_test.py | 22 +- tzrec/tests/prompt_test_util.py | 5 +- 11 files changed, 139 insertions(+), 427 deletions(-) diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index c12f406d..11972bba 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -40,7 +40,7 @@ remove_nullable, ) from tzrec.features.feature import BaseFeature -from tzrec.prompt.assembler import OUTPUT_KEYS, PROMPT_INFO_PREFIX, PromptAssembler +from tzrec.prompt.assembler import OUTPUT_KEYS, PromptAssembler from tzrec.prompt.types import CompiledPrompt from tzrec.protos import data_pb2 from tzrec.utils import config_util @@ -401,9 +401,7 @@ def _build_batch(self, input_data: Dict[str, pa.Array]) -> Batch: if self._prompt_assembler is not None: streams = self._prompt_assembler(output_data) - batch.additional_infos.update( - {PROMPT_INFO_PREFIX + k: streams[k] for k in OUTPUT_KEYS} - ) + batch.additional_infos.update({k: streams[k] for k in OUTPUT_KEYS}) # Set checkpoint info on batch batch.checkpoint_info = checkpoint_info diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index d72471c7..037582ee 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -29,10 +29,10 @@ from tzrec.models.genrec_model import BaseGenRecModel from tzrec.modules.dynamic_beam import capped_beam_widths, dynamic_beam_search from tzrec.prompt.assembler import ( - PROMPT_CU_SEQLENS, - PROMPT_INPUT_IDS, - PROMPT_MAX_SEQLEN, - PROMPT_RESPONSE_LENGTHS, + CU_SEQLENS, + INPUT_IDS, + MAX_SEQLEN, + RESPONSE_LENGTHS, ) from tzrec.prompt.types import CompiledPrompt from tzrec.protos.model_pb2 import ModelConfig @@ -189,8 +189,8 @@ def _left_pad_packed_inputs( Padded embeddings, attention mask and optional labels. """ infos = batch.additional_infos - cu_seqlens = infos[PROMPT_CU_SEQLENS] - max_seqlen = int(infos[PROMPT_MAX_SEQLEN]) + cu_seqlens = infos[CU_SEQLENS] + max_seqlen = int(infos[MAX_SEQLEN]) starts = cu_seqlens[:-1] lengths = cu_seqlens[1:] - starts batch_size = lengths.numel() @@ -205,8 +205,8 @@ def _left_pad_packed_inputs( if not build_labels: return padded, mask.long(), None - input_ids = infos[PROMPT_INPUT_IDS] - response_lengths = infos[PROMPT_RESPONSE_LENGTHS] + input_ids = infos[INPUT_IDS] + response_lengths = infos[RESPONSE_LENGTHS] labels = torch.full( (batch_size, max_seqlen), self._ignore_index, diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index b1cce726..5bc3f5da 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -18,10 +18,10 @@ from tzrec.datasets.utils import Batch from tzrec.models.genrec_causal_lm_model import GenRecCausalLMModel from tzrec.prompt.assembler import ( - PROMPT_CU_SEQLENS, - PROMPT_INPUT_IDS, - PROMPT_MAX_SEQLEN, - PROMPT_RESPONSE_LENGTHS, + CU_SEQLENS, + INPUT_IDS, + MAX_SEQLEN, + RESPONSE_LENGTHS, ) from tzrec.tests.prompt_test_util import ( _CODEBOOK, @@ -42,10 +42,10 @@ def test_packs_rows_of_different_lengths(self) -> None: ignore = -7 batch = Batch( additional_infos={ - PROMPT_CU_SEQLENS: cu, - PROMPT_INPUT_IDS: input_ids, - PROMPT_MAX_SEQLEN: torch.tensor(7), - PROMPT_RESPONSE_LENGTHS: response_lengths, + CU_SEQLENS: cu, + INPUT_IDS: input_ids, + MAX_SEQLEN: torch.tensor(7), + RESPONSE_LENGTHS: response_lengths, } ) model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 6fe3e9fb..ad91e7d2 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -37,12 +37,8 @@ HOLE_POSITIONS, HOLE_SLOT_COUNTS, INPUT_IDS, - PROMPT_CU_SEQLENS, - PROMPT_HOLE_POSITIONS, - PROMPT_HOLE_SLOT_COUNTS, - PROMPT_INPUT_IDS, ) -from tzrec.prompt.hole_keys import HOLE_KEYS, PROMPT_HOLE_KEYS +from tzrec.prompt.hole_keys import HOLE_KEYS from tzrec.prompt.types import CompiledPrompt, PromptPlan from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig @@ -209,7 +205,7 @@ def build_input(self, batch: Batch) -> torch.Tensor: Returns: ``(total_tokens, hidden_size)``. """ - ids = batch.additional_infos[PROMPT_INPUT_IDS] + ids = batch.additional_infos[INPUT_IDS] embeds = self.lm.get_input_embeddings()(ids) if not self._prompt.prompt_plan.projected_slots: return embeds @@ -222,7 +218,7 @@ def build_input(self, batch: Batch) -> torch.Tensor: ) # out of place: embeds carries grad from the embedding lookup return embeds.index_copy( - 0, batch.additional_infos[PROMPT_HOLE_POSITIONS], projected.to(embeds.dtype) + 0, batch.additional_infos[HOLE_POSITIONS], projected.to(embeds.dtype) ) def _tokens_to_local_codes( @@ -396,11 +392,11 @@ def predict(self, batch: Batch) -> Dict[str, torch.Tensor]: """ infos = batch.additional_infos out = { - INPUT_IDS: infos[PROMPT_INPUT_IDS], - CU_SEQLENS: infos[PROMPT_CU_SEQLENS], - HOLE_POSITIONS: infos[PROMPT_HOLE_POSITIONS], - HOLE_KEYS: infos[PROMPT_HOLE_KEYS], - HOLE_SLOT_COUNTS: infos[PROMPT_HOLE_SLOT_COUNTS], + INPUT_IDS: infos[INPUT_IDS], + CU_SEQLENS: infos[CU_SEQLENS], + HOLE_POSITIONS: infos[HOLE_POSITIONS], + HOLE_KEYS: infos[HOLE_KEYS], + HOLE_SLOT_COUNTS: infos[HOLE_SLOT_COUNTS], } if self._prompt.prompt_plan.projected_slots: out[SLOT_EMBEDS] = project_slots( @@ -412,6 +408,6 @@ def predict(self, batch: Batch) -> Dict[str, torch.Tensor]: ) else: out[SLOT_EMBEDS] = torch.zeros( - 0, 0, dtype=torch.float32, device=infos[PROMPT_INPUT_IDS].device + 0, 0, dtype=torch.float32, device=infos[INPUT_IDS].device ) return out diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 16d01d31..51ec8415 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -28,8 +28,6 @@ from tzrec.prompt.assembler import ( HOLE_POSITIONS, INPUT_IDS, - PROMPT_HOLE_POSITIONS, - PROMPT_INPUT_IDS, PromptAssembler, ) from tzrec.prompt.compile import compile_prompt @@ -129,8 +127,8 @@ def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: ) embeds = model.build_input(batch) - raw = model.lm.get_input_embeddings()(batch.additional_infos[PROMPT_INPUT_IDS]) - holes = batch.additional_infos[PROMPT_HOLE_POSITIONS] + raw = model.lm.get_input_embeddings()(batch.additional_infos[INPUT_IDS]) + holes = batch.additional_infos[HOLE_POSITIONS] self.assertGreater(holes.numel(), 0) changed = ~torch.isclose(embeds, raw).all(dim=-1) diff --git a/tzrec/models/model.py b/tzrec/models/model.py index 06bdb6a8..02636d2e 100644 --- a/tzrec/models/model.py +++ b/tzrec/models/model.py @@ -30,8 +30,8 @@ from tzrec.features.feature import BaseFeature from tzrec.loss.pe_mtl_loss import ParetoEfficientMultiTaskLoss from tzrec.modules.utils import BaseModule -from tzrec.prompt.assembler import OUTPUT_KEYS, PROMPT_INFO_PREFIX, PromptAssembler -from tzrec.prompt.hole_keys import PROMPT_HOLE_KEYS, HoleKeyBuilder +from tzrec.prompt.assembler import OUTPUT_KEYS, PromptAssembler +from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.protos.loss_pb2 import LossConfig from tzrec.protos.model_pb2 import FeatureGroupConfig, ModelConfig from tzrec.utils import config_util @@ -438,11 +438,9 @@ def get_batch( batch = self._data_parser.to_batch(data) if self._prompt_assembler is not None: streams = self._prompt_assembler(data) - batch.additional_infos.update( - {PROMPT_INFO_PREFIX + k: streams[k] for k in OUTPUT_KEYS} - ) + batch.additional_infos.update({k: streams[k] for k in OUTPUT_KEYS}) if self._hole_keys is not None: - batch.additional_infos[PROMPT_HOLE_KEYS] = self._hole_keys(data) + batch.additional_infos[HOLE_KEYS] = self._hole_keys(data) batch = batch.to(device, non_blocking=True) return batch diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index eb0d58cc..b7c1e33e 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -50,15 +50,6 @@ MAX_SEQLEN, ) -# where the collator stores the streams on the batch -PROMPT_INFO_PREFIX = "prompt_" -PROMPT_INPUT_IDS = PROMPT_INFO_PREFIX + INPUT_IDS -PROMPT_CU_SEQLENS = PROMPT_INFO_PREFIX + CU_SEQLENS -PROMPT_HOLE_POSITIONS = PROMPT_INFO_PREFIX + HOLE_POSITIONS -PROMPT_HOLE_SLOT_COUNTS = PROMPT_INFO_PREFIX + HOLE_SLOT_COUNTS -PROMPT_MAX_SEQLEN = PROMPT_INFO_PREFIX + MAX_SEQLEN -PROMPT_RESPONSE_LENGTHS = PROMPT_INFO_PREFIX + RESPONSE_LENGTHS - @torch.jit.script def _exclusive_cumsum(values: torch.Tensor) -> torch.Tensor: @@ -103,9 +94,8 @@ class PromptAssembler(nn.Module): """Walks a compiled plan to build one batch's packed token stream. The collator calls it eagerly on the host; export scripts the same module - into the serving front-end. Validation runs at both call sites: an assembled - row is checked, never truncated or repaired, because a stream that is - silently wrong reaches the loss or the beam as plausible output. + into the serving front-end. It trusts the parsed batch as every scripted + module does: the data contract belongs to the parser and to the request. Args: prompt_plan: the compiled walk order and its constants. @@ -120,14 +110,10 @@ class PromptAssembler(nn.Module): KIND_PROJECTED: Final[int] = 2 kinds: List[int] - names: List[str] static_tokens: List[List[int]] - exact_widths: List[int] hole_slots: List[int] is_sequences: List[bool] member_names: List[List[str]] - level_lo: List[int] - level_hi: List[int] def __init__( self, @@ -147,58 +133,37 @@ def __init__( if sid_space.sentinel_token_id is not None: self.sentinel = int(sid_space.sentinel_token_id) self.id_shift = int(sid_space.base_vocab_size) - self.num_levels = int(sid_space.num_levels) - self.level_lo = [int(o) for o in sid_space.level_offsets] - self.level_hi = [ - int(o + c) for o, c in zip(sid_space.level_offsets, sid_space.codebook) - ] - self.max_length = int(prompt_plan.max_length) self.kinds = [] - self.names = [] self.static_tokens = [] - self.exact_widths = [] self.hole_slots = [] self.is_sequences = [] self.member_names = [] + # the first slot member sizes the batch; an all-static plan reads the + # batch_size the parser passes along + self.anchor = "" # holes are grouped by projected occurrence in emission order, which is # the order of ``projected_slots``, of the front-end's projections and # of ``hole_keys`` occurrences = 0 - for index, seg in enumerate(segments): + for seg in segments: if isinstance(seg, Static): self._append( - self.KIND_STATIC, - "", - [int(t) for t in seg.token_ids], - -1, - -1, - False, - [], + self.KIND_STATIC, [int(t) for t in seg.token_ids], -1, False, [] ) continue assert isinstance(seg, SlotSeg) + if not self.anchor: + self.anchor = seg.feature_names[0] is_sequence = seg.group_type == FeatureGroupType.JAGGED_SEQUENCE if seg.fill is FillMode.INLINE: - # the answer's width sizes the loss window, so it is exact - width = -1 - if index >= self.num_body and seg.width.num_positions is not None: - width = int(seg.width.num_positions) self._append( - self.KIND_INLINE, - seg.name, - [], - width, - -1, - is_sequence, - [seg.feature_names[0]], + self.KIND_INLINE, [], -1, is_sequence, [seg.feature_names[0]] ) else: self._append( self.KIND_PROJECTED, - seg.name, [], - -1, occurrences, is_sequence, list(seg.feature_names), @@ -211,75 +176,32 @@ def __init__( def _append( self, kind: int, - name: str, tokens: List[int], - exact_width: int, hole_slot: int, is_sequence: bool, members: List[str], ) -> None: """Record one unrolled segment's constants.""" self.kinds.append(kind) - self.names.append(name) self.static_tokens.append(tokens) - self.exact_widths.append(exact_width) self.hole_slots.append(hole_slot) self.is_sequences.append(is_sequence) self.member_names.append(members) - def _lengths( - self, batch: Dict[str, torch.Tensor], slot: str, member: str, is_sequence: bool - ) -> torch.Tensor: - """Per-row item count of one member, as the data parser emits it.""" - key = member + ".lengths" - if key in batch: - return batch[key].to(torch.int64) - if is_sequence: - raise ValueError( - "prompt slot [" - + slot - + "] renders a sequence but the batch has no [" - + key - + "]: the column must be list, and under distributed " - + "embedding the processor must pass the raw parsed features " - + "through beside the looked-up embeddings." - ) - # a dense member has one row per sample and no lengths - return torch.ones( - batch[member + ".values"].size(0), - dtype=torch.int64, - device=batch_device(batch), - ) - def _batch_size(self, batch: Dict[str, torch.Tensor]) -> int: - """Row count, which every slot must agree on.""" - batch_size = -1 - for i in range(self.num_segments): - if self.kinds[i] == self.KIND_STATIC: - continue - for member in self.member_names[i]: - rows = int( - self._lengths( - batch, self.names[i], member, self.is_sequences[i] - ).numel() - ) - if batch_size < 0: - batch_size = rows - elif rows != batch_size: - raise ValueError( - "prompt slot [" - + self.names[i] - + "] has " - + str(rows) - + " samples, expected " - + str(batch_size) - + "." - ) - if batch_size >= 0: - return batch_size - if "batch_size" in batch: - return int(batch["batch_size"]) - return 0 + """Row count, from the anchor member. + + A sequence or a multi-value member carries ``lengths``, one per row; a + dense member has one row per sample and no lengths. + """ + if self.anchor == "": + if "batch_size" in batch: + return int(batch["batch_size"]) + return 0 + key = self.anchor + ".lengths" + if key in batch: + return int(batch[key].numel()) + return int(batch[self.anchor + ".values"].size(0)) def _inline_counts( self, batch: Dict[str, torch.Tensor], index: int, batch_size: int @@ -291,9 +213,7 @@ def _inline_counts( count is a segmented sum rather than ``lengths`` itself. """ member = self.member_names[index][0] - lengths = self._lengths( - batch, self.names[index], member, self.is_sequences[index] - ) + lengths = batch[member + ".lengths"].to(torch.int64) key = member + ".key_lengths" if key in batch: key_lengths = batch[key].to(torch.int64).reshape(-1) @@ -301,62 +221,6 @@ def _inline_counts( lengths = counts.index_add_(0, _row_ids(lengths), key_lengths) return lengths - def _inline_values( - self, batch: Dict[str, torch.Tensor], index: int, counts: torch.Tensor - ) -> torch.Tensor: - """Validate one INLINE segment's offset codes and shift them to token ids. - - The data carries ``level_offsets[l] + code``; the LM vocabulary needs one - further uniform shift by ``base_vocab_size``. - """ - name = self.names[index] - width = self.exact_widths[index] - if width >= 0: - wrong = torch.nonzero(counts != width) - if wrong.numel() > 0: - sample = int(wrong[0, 0]) - raise ValueError( - "prompt slot [" - + name - + "]: sample " - + str(sample) - + " has " - + str(int(counts[sample])) - + " values, but the compiled width is " - + str(width) - + ". The loss window is sized from that width, so a wider row " - + "would be supervised only in part." - ) - partial = torch.nonzero(counts % self.num_levels != 0) - if partial.numel() > 0: - sample = int(partial[0, 0]) - raise ValueError( - "prompt slot [" - + name - + "]: sample " - + str(sample) - + " has " - + str(int(counts[sample])) - + " values, not a whole number of " - + str(self.num_levels) - + "-level items." - ) - values = ( - batch[self.member_names[index][0] + ".values"].to(torch.int64).reshape(-1) - ) - by_level = values.reshape(-1, self.num_levels) - lo = torch.tensor(self.level_lo, dtype=torch.int64, device=values.device) - hi = torch.tensor(self.level_hi, dtype=torch.int64, device=values.device) - if bool(torch.any(by_level < lo)) or bool(torch.any(by_level >= hi)): - raise ValueError( - "prompt slot [" - + name - + "]: SID values must already carry their level offset, so level l " - + "lies in [level_offsets[l], level_offsets[l] + codebook[l]). Read " - + "the offset_codebook column, not codebook or origin_codebook." - ) - return values + self.id_shift - def _segment( self, batch: Dict[str, torch.Tensor], @@ -375,26 +239,16 @@ def _segment( seg_len = torch.full((batch_size,), width, dtype=torch.int64, device=device) return seg_len, run.unsqueeze(0).expand(batch_size, width).reshape(-1) + member = self.member_names[index][0] if kind == self.KIND_INLINE: + # the data carries ``level_offsets[l] + code``; the LM vocabulary + # needs one further uniform shift by ``base_vocab_size`` counts = self._inline_counts(batch, index, batch_size) - return counts, self._inline_values(batch, index, counts) + values = batch[member + ".values"].to(torch.int64).reshape(-1) + return counts, values + self.id_shift - members = self.member_names[index] - name = self.names[index] - seg_len = self._lengths(batch, name, members[0], self.is_sequences[index]) if self.is_sequences[index]: - for member in members[1:]: - other = self._lengths(batch, name, member, True) - if not torch.equal(seg_len, other): - raise ValueError( - "prompt slot [" - + name - + "] PROJECTED features [" - + members[0] - + "] and [" - + member - + "] have different per-sample lengths." - ) + seg_len = batch[member + ".lengths"].to(torch.int64) else: seg_len = torch.ones(batch_size, dtype=torch.int64, device=device) total = int(torch.sum(seg_len)) @@ -428,19 +282,6 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: stacked = torch.stack(seg_lens, dim=0) row_total = torch.sum(stacked, dim=0) - if self.max_length > 0: - over = torch.nonzero(row_total > self.max_length) - if over.numel() > 0: - sample = int(over[0, 0]) - raise ValueError( - "assembled sample " - + str(sample) - + " is " - + str(int(row_total[sample])) - + " tokens, over max_length " - + str(self.max_length) - + ". Samples are never truncated: cap the source features instead." - ) row_start = _exclusive_cumsum(row_total) seg_offsets = torch.cumsum(stacked, dim=0) - stacked diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index 02f1be84..3e30e333 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -13,9 +13,7 @@ import numpy as np import torch -from parameterized import parameterized -from tzrec.prompt import assembler from tzrec.prompt.assembler import ( CU_SEQLENS, HOLE_POSITIONS, @@ -35,8 +33,6 @@ WidthKind, ) from tzrec.protos.model_pb2 import FeatureGroupType -from tzrec.tests.prompt_test_util import assemble_into -from tzrec.utils.test_util import parameterized_name_func _BASE_VOCAB_SIZE = 1000 _SENTINEL = 1099 @@ -88,7 +84,7 @@ def _slot( ) -def _plan(segments, response=(), max_length=0) -> PromptPlan: +def _plan(segments, response=()) -> PromptPlan: projected = tuple( s for s in segments + tuple(response) @@ -97,7 +93,7 @@ def _plan(segments, response=(), max_length=0) -> PromptPlan: return PromptPlan( segments=tuple(segments), response_segments=tuple(response), - max_length=max_length, + max_length=0, max_total_length=None, max_holes=0, logits_suffix_len=None, @@ -106,10 +102,10 @@ def _plan(segments, response=(), max_length=0) -> PromptPlan: ) -def _asm(segments, response=(), max_length=0, sid_space=None) -> PromptAssembler: +def _asm(segments, response=(), sid_space=None) -> PromptAssembler: """An assembler over one ad-hoc plan.""" return PromptAssembler( - _plan(segments, response=response, max_length=max_length), + _plan(segments, response=response), _sid_space() if sid_space is None else sid_space, ) @@ -226,22 +222,6 @@ def test_response_is_optional_and_its_length_is_recorded(self) -> None: ) self.assertEqual(prompt_only[RESPONSE_LENGTHS].tolist(), [0]) - def test_rejects_a_response_of_the_wrong_width(self) -> None: - plan = _plan( - (_slot("hist", FillMode.INLINE),), - response=(_slot("answer", FillMode.INLINE),), - ) - asm = PromptAssembler(plan, _sid_space()) - with self.assertRaisesRegex(ValueError, "compiled width is 3"): - asm( - _parsed( - { - "hist": [np.array([1, 6, 11])], - "answer": [np.array([0, 4, 8, 1, 5, 9])], - } - ) - ) - def test_a_multi_value_history_walks_like_the_flat_layout(self) -> None: """Items with key_lengths and one code per position are one stream.""" asm = _asm((Static((7,)), _slot("hist", FillMode.INLINE))) @@ -258,94 +238,34 @@ def test_a_multi_value_history_walks_like_the_flat_layout(self) -> None: for key in (INPUT_IDS, CU_SEQLENS, HOLE_POSITIONS, MAX_SEQLEN): self.assertTrue(torch.equal(flat[key], items[key]), key) - def test_scalar_label_is_named_not_a_key_error(self) -> None: - asm = _asm((_slot("hist", FillMode.INLINE),)) - - with self.assertRaisesRegex(ValueError, "must be\\s+list"): - asm({"hist.values": torch.tensor([1, 6, 11])}) - - @parameterized.expand( - [ - [[1, 2, 3]], - [[1, 6, 12]], - ], - name_func=parameterized_name_func, - ) - def test_rejects_a_code_outside_its_band(self, values) -> None: - asm = _asm((_slot("hist", FillMode.INLINE),)) - with self.assertRaisesRegex(ValueError, "offset_codebook column"): - asm(_parsed({"hist": [np.array(values)]})) - - def test_rejects_a_partial_item(self) -> None: - asm = _asm((_slot("hist", FillMode.INLINE),)) - with self.assertRaisesRegex(ValueError, "whole number of 3-level items"): - asm(_parsed({"hist": [np.array([1, 6])]})) - - def test_over_long_row_is_an_error_not_a_truncation(self) -> None: - asm = _asm((Static((7, 8, 9)), _slot("hist", FillMode.INLINE)), max_length=4) - with self.assertRaisesRegex(ValueError, "never truncated"): - asm(_parsed({"hist": [np.array([1, 6, 11])]})) - def test_column_shaped_values_are_flattened(self) -> None: # the data parser emits (total, value_dim) for a dense sequence feature - from tzrec.prompt.types import CompiledPrompt, ProjectionPlan - - plan = _plan((_slot("hist", FillMode.INLINE),)) - compiled_prompt = CompiledPrompt( - sid_space=_sid_space(), - prompt_plan=plan, - projection_plan=ProjectionPlan(projections={}, slot_to_module={}), + asm = _asm((_slot("hist", FillMode.INLINE),)) + out = asm( + { + "hist.values": torch.tensor([[1], [6], [11], [0], [4], [8]]), + "hist.lengths": torch.tensor([3, 3]), + } ) - parsed = { - "hist.values": np.array([[1], [6], [11], [0], [4], [8]]), - "hist.lengths": np.array([3, 3]), - } - out = assemble_into(compiled_prompt, parsed) - self.assertEqual(out["prompt_cu_seqlens"].tolist(), [0, 3, 6]) - self.assertEqual(out["prompt_input_ids"].tolist()[0], _BASE_VOCAB_SIZE + 1) - self.assertEqual(int(out["prompt_max_seqlen"]), 3) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 3, 6]) + self.assertEqual(out[INPUT_IDS].tolist()[0], _BASE_VOCAB_SIZE + 1) + self.assertEqual(int(out[MAX_SEQLEN]), 3) - def test_rejects_inconsistent_slot_batch_sizes(self) -> None: - plan = _plan( + def test_the_first_slot_member_sizes_the_batch(self) -> None: + """A dense anchor counts rows; a jagged anchor counts lengths.""" + dense = _asm( ( - _slot("hist", FillMode.INLINE), - _slot("answer", FillMode.INLINE), + Static((7,)), + _slot("vec", FillMode.PROJECTED, group_type=FeatureGroupType.DEEP), ) ) - asm = PromptAssembler(plan, _sid_space()) - parsed = { - "hist.values": torch.tensor([1, 6, 11]), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor([0, 4, 8, 1, 6, 11]), - "answer.lengths": torch.tensor([3, 3]), - } - with self.assertRaisesRegex( - ValueError, r"prompt slot \[answer\] has 2 samples, expected 1" - ): - asm(parsed) - - def test_rejects_mismatched_projected_member_lengths(self) -> None: - plan = _plan( - ( - _slot( - "profile", - FillMode.PROJECTED, - 4, - feature_names=("age", "country"), - ), - ) + out = dense({"vec.values": torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])}) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 2, 4, 6]) + jagged = _asm((Static((7,)), _slot("beh", FillMode.PROJECTED, 4))) + out = jagged( + _parsed(projected={"beh": [1, 2]}) | {"beh.values": torch.tensor([1, 2, 3])} ) - asm = PromptAssembler(plan, _sid_space()) - parsed = { - "age.lengths": torch.tensor([2, 1]), - "country.lengths": torch.tensor([2, 2]), - } - - with self.assertRaisesRegex( - ValueError, - r"PROJECTED features \[age\] and \[country\] have different", - ): - asm(parsed) + self.assertEqual(out[CU_SEQLENS].tolist(), [0, 2, 5]) def test_deep_projected_members_emit_one_hole_per_sample(self) -> None: plan = _plan( @@ -387,7 +307,6 @@ def test_output_keys_match_the_module_constants(self) -> None: }, ) self.assertEqual(out[INPUT_IDS].tolist(), [7, 7]) - self.assertEqual(assembler.PROMPT_INPUT_IDS, "prompt_" + INPUT_IDS) def test_scripting_preserves_every_output(self) -> None: """The artifact and the collator's module are the same function.""" @@ -414,14 +333,6 @@ def test_scripting_preserves_every_output(self) -> None: for key, value in scripted(batch).items(): self.assertTrue(torch.equal(eager[key], value), key) - def test_a_scripted_walk_reports_its_validation(self) -> None: - """The exported artifact refuses a bad request with the same message.""" - scripted = torch.jit.script( - _asm((Static((7, 8, 9)), _slot("hist", FillMode.INLINE)), max_length=4) - ) - with self.assertRaisesRegex(torch.jit.Error, "never truncated"): - scripted(_parsed({"hist": [np.array([1, 6, 11])]})) - if __name__ == "__main__": unittest.main() diff --git a/tzrec/prompt/hole_keys.py b/tzrec/prompt/hole_keys.py index 05830f32..bf0942ff 100644 --- a/tzrec/prompt/hole_keys.py +++ b/tzrec/prompt/hole_keys.py @@ -25,17 +25,11 @@ import torch from torch import nn -from tzrec.prompt.assembler import ( - PROMPT_INFO_PREFIX, - _row_ids, - _within_row_index, - batch_device, -) +from tzrec.prompt.assembler import _row_ids, _within_row_index, batch_device from tzrec.prompt.types import PromptPlan from tzrec.protos.model_pb2 import FeatureGroupType HOLE_KEYS = "hole_keys" -PROMPT_HOLE_KEYS = PROMPT_INFO_PREFIX + HOLE_KEYS @torch.jit.script @@ -71,7 +65,8 @@ class HoleKeyBuilder(nn.Module): slots holding the same id would match; without the last two, a two-member slot with values ``(a, b)`` would match one with ``(b, a)`` and a permuted multi-value item would match itself reordered -- all plausible, all wrong, - and all silent. + and all silent. A dense member contributes its float32 bit pattern per row, + which is the parsed input and not a computed reduction. Args: prompt_plan: the compiled plan; its ``projected_slots`` fix the hole @@ -87,94 +82,66 @@ class HoleKeyBuilder(nn.Module): # enough that a position cannot carry into the member index MEMBER_STRIDE: Final[int] = 1 << 32 - names: List[str] member_names: List[List[str]] is_sequences: List[bool] salts: List[int] def __init__(self, prompt_plan: PromptPlan) -> None: super().__init__() - self.names = [] self.member_names = [] self.is_sequences = [] self.salts = [] for seg in prompt_plan.projected_slots: - self.names.append(seg.name) self.member_names.append(list(seg.feature_names)) self.is_sequences.append(seg.group_type == FeatureGroupType.JAGGED_SEQUENCE) self.salts.append(_wrap64(self.C_SLOT * int(seg.slot_id))) - self.num_slots = len(self.names) - - def _lengths(self, batch: Dict[str, torch.Tensor], member: str) -> torch.Tensor: - """Per-row item count; a dense member has one row per sample and no lengths.""" - key = member + ".lengths" - if key in batch: - return batch[key].to(torch.int64) - return torch.ones( - batch[member + ".values"].size(0), - dtype=torch.int64, - device=batch_device(batch), - ) + self.num_slots = len(self.member_names) def _fold_slot(self, batch: Dict[str, torch.Tensor], index: int) -> torch.Tensor: """One projected occurrence's keys, one per hole, in sample order.""" salt = self.salts[index] - name = self.names[index] members = self.member_names[index] is_sequence = self.is_sequences[index] # the assembler's hole count: one per item of a sequence slot, one per - # sample of a DEEP slot - first = self._lengths(batch, members[0]) - num_holes = int(torch.sum(first)) if is_sequence else int(first.numel()) + # sample of a DEEP slot; a dense member has one row per hole either way + first = batch[members[0] + ".values"] + if first.is_floating_point(): + num_holes = int(first.size(0)) + else: + lengths = batch[members[0] + ".lengths"].to(torch.int64) + num_holes = int(torch.sum(lengths)) if is_sequence else int(lengths.numel()) keys = torch.zeros(num_holes, dtype=torch.int64, device=first.device) for member_index in range(len(members)): member = members[member_index] raw = batch[member + ".values"] - key_length_key = member + ".key_lengths" - - if not is_sequence: - # a dense member contributes its float32 bit pattern verbatim, - # which is the parsed input and not a computed reduction, so - # it is stable for a given request - if raw.is_floating_point(): - width = raw.size(1) - values = ( - raw.to(torch.float32) - .contiguous() - .view(torch.int32) - .to(torch.int64) - .reshape(-1) - & 0xFFFFFFFF - ) - hole = torch.repeat_interleave( - torch.arange(num_holes, dtype=torch.int64, device=raw.device), - torch.full( - (num_holes,), width, dtype=torch.int64, device=raw.device - ), - ) - local = ( - torch.arange(width, dtype=torch.int64, device=raw.device) - .unsqueeze(0) - .expand(num_holes, width) - .reshape(-1) - ) - else: - lengths = self._lengths(batch, member) - values = raw.to(torch.int64).reshape(-1) - hole = _row_ids(lengths) - local = _within_row_index(lengths) + if raw.is_floating_point(): + rows = raw.size(0) + width = raw.size(1) + values = ( + raw.to(torch.float32) + .contiguous() + .view(torch.int32) + .to(torch.int64) + .reshape(-1) + & 0xFFFFFFFF + ) + hole = torch.arange( + rows, dtype=torch.int64, device=raw.device + ).repeat_interleave(width) + local = ( + torch.arange(width, dtype=torch.int64, device=raw.device) + .unsqueeze(0) + .expand(rows, width) + .reshape(-1) + ) else: - if raw.is_floating_point(): - raise ValueError( - "prompt slot [" - + name - + "] member [" - + member - + "] is a dense sequence; the fold has no per-item boundary " - + "for it." - ) values = raw.to(torch.int64).reshape(-1) - if key_length_key in batch: + key_length_key = member + ".key_lengths" + if not is_sequence: + lengths = batch[member + ".lengths"].to(torch.int64) + hole = _row_ids(lengths) + local = _within_row_index(lengths) + elif key_length_key in batch: key_lengths = batch[key_length_key].to(torch.int64).reshape(-1) hole = _row_ids(key_lengths) local = _within_row_index(key_lengths) diff --git a/tzrec/prompt/hole_keys_test.py b/tzrec/prompt/hole_keys_test.py index 22bb3cdc..717cb7c5 100644 --- a/tzrec/prompt/hole_keys_test.py +++ b/tzrec/prompt/hole_keys_test.py @@ -16,7 +16,7 @@ from torch import nn from tzrec.prompt.assembler import HOLE_SLOT_COUNTS, PromptAssembler -from tzrec.prompt.hole_keys import PROMPT_HOLE_KEYS, HoleKeyBuilder, mix64 +from tzrec.prompt.hole_keys import HoleKeyBuilder, mix64 from tzrec.prompt.types import ( FillMode, PromptPlan, @@ -192,20 +192,24 @@ def keys(rows): self.assertNotEqual(keys([[0.5, 1.0]]).tolist(), keys([[1.0, 0.5]]).tolist()) self.assertEqual(keys([[0.5, 1.0], [0.5, 1.0]]).numel(), 2) - def test_a_dense_sequence_member_is_rejected(self) -> None: - with self.assertRaisesRegex(ValueError, "no per-item boundary"): - self.module( - { - "beh.values": torch.tensor([[0.5], [1.0]]), - "beh.lengths": torch.tensor([2]), - } + def test_a_dense_sequence_member_folds_each_item(self) -> None: + """Every row of a dense sequence is one item, so one hole and one key.""" + + def keys(rows): + return self.module( + {"beh.values": torch.tensor(rows), "beh.lengths": torch.tensor([2])} ) + self.assertEqual(keys([[0.5], [1.0]]).numel(), 2) + self.assertEqual(keys([[0.5], [1.0]]).tolist(), keys([[0.5], [1.0]]).tolist()) + self.assertNotEqual( + keys([[0.5], [1.0]]).tolist(), keys([[1.0], [0.5]]).tolist() + ) + def test_no_projected_slot_folds_nothing(self) -> None: keys = HoleKeyBuilder(_plan((Static((7,)),)))({"batch_size": torch.tensor(2)}) self.assertEqual(keys.numel(), 0) self.assertEqual(keys.dtype, torch.int64) - self.assertEqual(PROMPT_HOLE_KEYS, "prompt_hole_keys") def test_scripting_preserves_the_keys(self) -> None: """The exported front-end folds exactly what the eager module does.""" diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py index 551f34d5..02e35cdb 100644 --- a/tzrec/tests/prompt_test_util.py +++ b/tzrec/tests/prompt_test_util.py @@ -21,7 +21,7 @@ from tzrec.datasets.utils import BASE_DATA_GROUP, Batch from tzrec.features.feature import BaseFeature, FgMode, create_features from tzrec.main import _create_model -from tzrec.prompt.assembler import PROMPT_INFO_PREFIX, PromptAssembler +from tzrec.prompt.assembler import PromptAssembler from tzrec.prompt.compile import compile_prompt from tzrec.prompt.types import CompiledPrompt from tzrec.protos import feature_pb2 @@ -93,10 +93,9 @@ def assemble_into( The assembled streams keyed for ``additional_infos``. """ batch = {k: torch.as_tensor(np.asarray(v)) for k, v in parsed_features.items()} - streams = PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( + return PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( batch ) - return {PROMPT_INFO_PREFIX + k: v for k, v in streams.items()} _CODEBOOK = [4, 4, 4] From 044d2bd6df56c97f4849fb6767793703a8365a0c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 14:48:23 +0800 Subject: [PATCH 13/21] [refactor] fold the genrec test fixtures into test_util create_genrec_test_tokenizer and create_genrec_test_model sit beside create_tiny_causal_lm; the genrec tests are plain TestCases again and tzrec/tests/prompt_test_util.py is gone. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/models/genrec_causal_lm_model_test.py | 59 +++-- tzrec/models/genrec_model_test.py | 267 +++++++++----------- tzrec/prompt/compile_test.py | 64 +++-- tzrec/tests/genrec_integration_test.py | 18 +- tzrec/tests/prompt_test_util.py | 163 ------------ tzrec/utils/test_util.py | 103 +++++++- 6 files changed, 295 insertions(+), 379 deletions(-) delete mode 100644 tzrec/tests/prompt_test_util.py diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 5bc3f5da..8c9da960 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -22,13 +22,13 @@ INPUT_IDS, MAX_SEQLEN, RESPONSE_LENGTHS, + PromptAssembler, ) -from tzrec.tests.prompt_test_util import ( - _CODEBOOK, - GenRecModelTestBase, - offset_sid_codes, +from tzrec.utils.test_util import ( + create_genrec_test_model, + make_test_dir, + parameterized_name_func, ) -from tzrec.utils.test_util import parameterized_name_func class LeftPadPackedInputsTest(unittest.TestCase): @@ -72,9 +72,12 @@ def test_packs_rows_of_different_lengths(self) -> None: ) -class GenRecCausalLMModelTest(GenRecModelTestBase): +class GenRecCausalLMModelTest(unittest.TestCase): """The decode schedule and the training forward, both subclass-owned.""" + def setUp(self) -> None: + self.test_dir = make_test_dir() + @parameterized.expand( [ [[2, 3, 4], [2, 3, 4]], @@ -84,45 +87,45 @@ class GenRecCausalLMModelTest(GenRecModelTestBase): name_func=parameterized_name_func, ) def test_beam_widths_are_capped_once_at_init(self, beam_widths, expected) -> None: - model = self._model(beam_widths=beam_widths, num_return_sequences=1) + model, compiled_prompt = create_genrec_test_model( + self.test_dir, beam_widths=beam_widths, num_return_sequences=1 + ) self.assertEqual(model._capped_widths, expected) - space = self.compiled_prompt.sid_space + space = compiled_prompt.sid_space self.assertEqual(model._bands, list(zip(space.band_lo, space.band_hi))) def test_rejects_a_schedule_that_does_not_match_the_codebook(self) -> None: with self.assertRaisesRegex(ValueError, "entries but the codebook has"): - self._model(beam_widths=(2, 2)) + create_genrec_test_model(self.test_dir, beam_widths=(2, 2)) def test_rejects_a_non_positive_beam_width(self) -> None: with self.assertRaisesRegex(ValueError, "must be >= 1"): - self._model(beam_widths=(2, 0, 2)) + create_genrec_test_model(self.test_dir, beam_widths=(2, 0, 2)) def test_beam_config_uses_final_capped_capacity(self) -> None: with self.assertRaisesRegex(ValueError, "final capped beam width \\(4\\)"): - self._model( - beam_widths=(1, 1, 100), - num_return_sequences=5, + create_genrec_test_model( + self.test_dir, beam_widths=(1, 1, 100), num_return_sequences=5 ) def test_training_forward_builds_no_cache(self) -> None: - model = self._model() + model, compiled_prompt = create_genrec_test_model(self.test_dir) + batch = Batch() + batch.additional_infos.update( + PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( + { + # offset SID codes for the (4, 4, 4) codebook + "hist.values": torch.tensor([0, 5, 10]), + "hist.lengths": torch.tensor([3]), + "answer.values": torch.tensor([1, 6, 11]), + "answer.lengths": torch.tensor([3]), + } + ) + ) inner = model.lm.model.forward with mock.patch.object(model.lm.model, "forward", side_effect=inner) as spy: - model.predict( - self._batch( - { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2], _CODEBOOK) - ), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor( - offset_sid_codes([1, 2, 3], _CODEBOOK) - ), - "answer.lengths": torch.tensor([3]), - } - ) - ) + model.predict(batch) self.assertIs(spy.call_args.kwargs["use_cache"], False) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 51ec8415..62088bd2 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -9,6 +9,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os import unittest import torch @@ -17,7 +18,8 @@ from torchrec import KeyedJaggedTensor from transformers import AutoModelForCausalLM -from tzrec.datasets.utils import Batch +from tzrec.datasets.utils import BASE_DATA_GROUP, Batch +from tzrec.main import _create_model from tzrec.models.genrec_model import ( _PARAM_DTYPE, SLOT_EMBEDS, @@ -25,36 +27,79 @@ project_slots, ) from tzrec.models.model import ScriptWrapper, TrainWrapper -from tzrec.prompt.assembler import ( - HOLE_POSITIONS, - INPUT_IDS, - PromptAssembler, -) -from tzrec.prompt.compile import compile_prompt +from tzrec.prompt.assembler import HOLE_POSITIONS, INPUT_IDS, PromptAssembler from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder +from tzrec.protos import feature_pb2 +from tzrec.protos.model_pb2 import ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -from tzrec.protos.prompt_pb2 import PromptConfig -from tzrec.tests.prompt_test_util import ( - _CODEBOOK, - _HIST, - GenRecModelTestBase, - assemble_into, - create_prompt_feature, - offset_sid_codes, - projected_feature, -) +from tzrec.protos.prompt_pb2 import PromptSlot from tzrec.utils.fx_util import symbolic_trace from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( + create_genrec_test_model, + make_test_dir, parameterized_name_func, ) +# offset SID codes for the (4, 4, 4) codebook: level_offsets[l] + code +_HIST_CODES = [0, 5, 10] +_LONG_HIST_CODES = [0, 5, 10, 3, 4, 9] +_ANSWER_CODES = [1, 6, 11] + + +def _hist() -> feature_pb2.FeatureConfig: + return feature_pb2.FeatureConfig( + sequence_raw_feature=feature_pb2.RawFeature( + feature_name="hist", expression="user:hist" + ) + ) + + +def _projected(name: str, dim: int) -> feature_pb2.FeatureConfig: + return feature_pb2.FeatureConfig( + sequence_id_feature=feature_pb2.IdFeature( + feature_name=name, + expression=f"user:{name}", + num_buckets=32, + embedding_dim=dim, + sequence_length=2, + ) + ) + -class BaseGenRecModelTest(GenRecModelTestBase): +def _batch(compiled_prompt, parsed, sparse=None) -> Batch: + batch = Batch(sparse_features={BASE_DATA_GROUP: sparse} if sparse else {}) + batch.additional_infos.update( + PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)(parsed) + ) + return batch + + +def _projected_batch(compiled_prompt) -> Batch: + return _batch( + compiled_prompt, + { + "hist.values": torch.tensor(_HIST_CODES).reshape(-1, 1), + "hist.lengths": torch.tensor([3]), + "answer.values": torch.tensor(_ANSWER_CODES), + "answer.lengths": torch.tensor([3]), + "prof.values": torch.tensor([5, 9]), + "prof.lengths": torch.tensor([2]), + }, + sparse=KeyedJaggedTensor.from_lengths_sync( + keys=["prof"], values=torch.tensor([5, 9]), lengths=torch.tensor([2]) + ), + ) + + +class BaseGenRecModelTest(unittest.TestCase): """Shared causal-LM behavior, reached through its concrete subclass.""" + def setUp(self) -> None: + self.test_dir = make_test_dir() + self.model, self.compiled_prompt = create_genrec_test_model(self.test_dir) + def test_tokens_to_local_codes_undoes_shifts_and_groups_beams(self) -> None: - model = self._model() space = self.compiled_prompt.sid_space local_codes = torch.tensor( [ @@ -65,66 +110,44 @@ def test_tokens_to_local_codes_undoes_shifts_and_groups_beams(self) -> None: ] ) tokens = local_codes + torch.tensor(space.level_offsets) + space.base_vocab_size - codes = model._tokens_to_local_codes(tokens, batch_size=2) + codes = self.model._tokens_to_local_codes(tokens, batch_size=2) self.assertEqual(codes.shape, (2, 2, space.num_levels)) self.assertEqual(codes.tolist(), local_codes.reshape(2, 2, -1).tolist()) def test_rejects_a_model_built_without_a_prompt(self) -> None: + model_config = ModelConfig() + model_config.genrec_causal_lm_model.hf_model_name_or_path = os.path.join( + self.test_dir, "backbone" + ) with self.assertRaisesRegex(ValueError, "needs a compiled prompt"): - self._model(compiled_prompt=None) + _create_model(model_config, [], ["answer"], compiled_prompt=None) def test_shared_projection_name_requires_matching_widths(self) -> None: - features = [ - create_prompt_feature(_HIST), - create_prompt_feature(projected_feature("pa", 8)), - create_prompt_feature(projected_feature("pb", 16)), - ] - cfg = PromptConfig( - tokenizer_path=self.tok, - prompt="History : {{hist}} . {{pa}} {{pb}} Predict :", - response="{{answer}}", - ) - cfg.sid_space.codebook.extend(_CODEBOOK) - for name in ("pa", "pb"): - slot = cfg.slots.add(name=name, projection_name="shared") - slot.feature_names.append(name) - compiled_prompt = compile_prompt(cfg, features, ["answer"]) - with self.assertRaisesRegex(ValueError, "cannot share a module"): - self._model(features=features, compiled_prompt=compiled_prompt) + create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("pa", 8), _projected("pb", 16)], + prompt="History : {{hist}} . {{pa}} {{pb}} Predict :", + slots=[ + PromptSlot( + name="pa", feature_names=["pa"], projection_name="shared" + ), + PromptSlot( + name="pb", feature_names=["pb"], projection_name="shared" + ), + ], + ) def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: - features = [ - create_prompt_feature(_HIST), - create_prompt_feature(projected_feature("prof", 8)), - ] - compiled_prompt = self._compile( - features, - template="History : {{hist}} . Predict {{prof}} :", - response="{{answer}}", + model, compiled_prompt = create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("prof", 8)], + prompt="History : {{hist}} . Predict {{prof}} :", ) - model = self._model(features=features, compiled_prompt=compiled_prompt) # the embedding table is built on meta until something materializes it init_parameters(model, device=torch.device("cpu")) - batch = self._batch( - { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2], _CODEBOOK) - ).reshape(-1, 1), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor(offset_sid_codes([1, 2, 3], _CODEBOOK)), - "answer.lengths": torch.tensor([3]), - "prof.values": torch.tensor([5, 9]), - "prof.lengths": torch.tensor([2]), - }, - compiled_prompt=compiled_prompt, - sparse=KeyedJaggedTensor.from_lengths_sync( - keys=["prof"], - values=torch.tensor([5, 9]), - lengths=torch.tensor([2]), - ), - ) + batch = _projected_batch(compiled_prompt) embeds = model.build_input(batch) raw = model.lm.get_input_embeddings()(batch.additional_infos[INPUT_IDS]) @@ -143,39 +166,14 @@ def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: name_func=parameterized_name_func, ) def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> None: - features = [ - create_prompt_feature(_HIST), - create_prompt_feature(projected_feature("prof", 8)), - ] - compiled_prompt = self._compile( - features, - template="History : {{hist}} . Predict {{prof}} :", - response="{{answer}}", - ) - model = self._model( - features=features, - compiled_prompt=compiled_prompt, + model, compiled_prompt = create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("prof", 8)], + prompt="History : {{hist}} . Predict {{prof}} :", lm_parameter_dtype=lm_parameter_dtype, ) init_parameters(model, device=torch.device("cpu")) - batch = self._batch( - { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2], _CODEBOOK) - ).reshape(-1, 1), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor(offset_sid_codes([1, 2, 3], _CODEBOOK)), - "answer.lengths": torch.tensor([3]), - "prof.values": torch.tensor([5, 9]), - "prof.lengths": torch.tensor([2]), - }, - compiled_prompt=compiled_prompt, - sparse=KeyedJaggedTensor.from_lengths_sync( - keys=["prof"], - values=torch.tensor([5, 9]), - lengths=torch.tensor([2]), - ), - ) + batch = _projected_batch(compiled_prompt) embeds = model.build_input(batch) self.assertIs(embeds.dtype, _PARAM_DTYPE[lm_parameter_dtype]) @@ -190,89 +188,74 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N self.assertGreater(float(proj.head.weight.grad.abs().sum()), 0.0) def test_metric_averages_the_loss_across_batches(self) -> None: - model = self._model() - model.init_metric() + self.model.init_metric() for value in (1.0, 3.0): - model.update_metric({}, Batch(), {"ce_loss": torch.tensor(value)}) + self.model.update_metric({}, Batch(), {"ce_loss": torch.tensor(value)}) self.assertAlmostEqual( - model._metric_modules["ce_loss"].compute().item(), 2.0, places=5 + self.model._metric_modules["ce_loss"].compute().item(), 2.0, places=5 ) def test_init_from_pretrained_replaces_the_empty_weights(self) -> None: - model = self._model() base_vocab_size = self.compiled_prompt.sid_space.base_vocab_size - before = model.lm.get_input_embeddings().weight[:base_vocab_size].clone() - model.init_from_pretrained() - after = model.lm.get_input_embeddings().weight[:base_vocab_size] + embeddings = self.model.lm.get_input_embeddings() + before = embeddings.weight[:base_vocab_size].clone() + self.model.init_from_pretrained() + after = embeddings.weight[:base_vocab_size] # the checkpoint rows land verbatim; only the appended SID rows are new - reference = AutoModelForCausalLM.from_pretrained(self.backbone) + reference = AutoModelForCausalLM.from_pretrained( + os.path.join(self.test_dir, "backbone") + ) expected = reference.get_input_embeddings().weight[:base_vocab_size] self.assertFalse(torch.allclose(before, expected)) torch.testing.assert_close(after, expected) - def _batch_from_codes(self, hist, answer): - parsed = { - "hist.values": torch.tensor(offset_sid_codes(hist, _CODEBOOK)), - "hist.lengths": torch.tensor([len(hist)]), - "answer.values": torch.tensor(offset_sid_codes(answer, _CODEBOOK)), - "answer.lengths": torch.tensor([len(answer)]), - } - batch = Batch() - batch.additional_infos.update(assemble_into(self.compiled_prompt, parsed)) - return batch - def test_model_resizes_to_target_vocab_size(self) -> None: - model = self._model() - rows = model.lm.get_input_embeddings().weight.shape[0] + rows = self.model.lm.get_input_embeddings().weight.shape[0] self.assertEqual(rows, self.compiled_prompt.sid_space.target_vocab_size) self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: - model = self._model() - batch = self._batch_from_codes([0, 1, 2, 3, 0, 1], [1, 2, 3]) - predictions = model.predict(batch) - loss = model.loss(predictions, batch)["ce_loss"] + batch = _batch( + self.compiled_prompt, + { + "hist.values": torch.tensor(_LONG_HIST_CODES), + "hist.lengths": torch.tensor([6]), + "answer.values": torch.tensor(_ANSWER_CODES), + "answer.lengths": torch.tensor([3]), + }, + ) + predictions = self.model.predict(batch) + loss = self.model.loss(predictions, batch)["ce_loss"] self.assertTrue(bool(torch.isfinite(loss))) loss.backward() - grad = model.lm.get_input_embeddings().weight.grad + grad = self.model.lm.get_input_embeddings().weight.grad self.assertIsNotNone(grad) self.assertTrue(bool((grad.abs().sum() > 0))) def test_training_forward_survives_fx_tracing(self) -> None: - model = self._model() - - torch.fx.symbolic_trace(TrainWrapper(model)) + torch.fx.symbolic_trace(TrainWrapper(self.model)) -class GenRecFrontEndTest(GenRecModelTestBase): +class GenRecFrontEndTest(unittest.TestCase): """The served half of the model, under the same wrapper every export uses.""" def setUp(self) -> None: - super().setUp() - self.features = [ - create_prompt_feature(_HIST), - create_prompt_feature(projected_feature("beh", 8)), - ] - self.prompt_config = PromptConfig( - tokenizer_path=self.tok, + self.test_dir = make_test_dir() + self.model, self.compiled_prompt = create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("beh", 8)], prompt="History : {{hist}} . {{beh}} Predict :", - response="{{answer}}", ) - self.prompt_config.sid_space.codebook.extend(_CODEBOOK) - self.compiled_prompt = compile_prompt( - self.prompt_config, self.features, ["answer"] - ) - self.model = self._model() init_parameters(self.model, device=torch.device("cpu")) # the parsed dict as the data parser emits it: a dense sequence of codes # and a sparse behaviour sequence self.data = { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2, 3, 0, 1], _CODEBOOK), dtype=torch.float32 - ).reshape(-1, 1), + "hist.values": torch.tensor(_LONG_HIST_CODES, dtype=torch.float32).reshape( + -1, 1 + ), "hist.lengths": torch.tensor([6]), "beh.values": torch.tensor([3, 9]), "beh.lengths": torch.tensor([2]), diff --git a/tzrec/prompt/compile_test.py b/tzrec/prompt/compile_test.py index 4017d694..9dbfad76 100644 --- a/tzrec/prompt/compile_test.py +++ b/tzrec/prompt/compile_test.py @@ -21,11 +21,7 @@ from tzrec.prompt.types import FillMode, SlotSeg, Static, WidthKind from tzrec.protos import feature_pb2 from tzrec.protos.prompt_pb2 import PromptConfig -from tzrec.tests.prompt_test_util import ( - create_prompt_feature, - create_prompt_tokenizer, -) -from tzrec.utils.test_util import make_test_dir +from tzrec.utils.test_util import create_genrec_test_tokenizer, make_test_dir _WORDS = ["History", "Profile", "Predict", ":", ".", "Histor0", "", "<|im_end|>"] @@ -38,10 +34,16 @@ _AGE = 'id_feature { feature_name: "age" expression: "user:age" num_buckets: 8 }' +def _feature(text: str): + config = feature_pb2.FeatureConfig() + text_format.Merge(text, config) + return create_features([config], fg_mode=FgMode.FG_NONE)[0] + + class CompilePromptTest(unittest.TestCase): def setUp(self) -> None: self.test_dir = make_test_dir() - self.tok_path = create_prompt_tokenizer( + self.tok_path = create_genrec_test_tokenizer( os.path.join(self.test_dir, "tok.json"), _WORDS ) @@ -56,7 +58,7 @@ def _compile(self, cfg, features): def test_sid_space_resolves_offsets_and_bands(self) -> None: cfg = self._config(prompt="History : {{hist}}") cfg.sid_space.codebook.extend([4, 4, 4]) - compiled = self._compile(cfg, [create_prompt_feature(_HIST)]) + compiled = self._compile(cfg, [_feature(_HIST)]) space = compiled.sid_space base_vocab_size = space.base_vocab_size @@ -78,9 +80,7 @@ def test_sid_space_resolves_offsets_and_bands(self) -> None: def test_inline_needs_no_group_projected_gets_one(self) -> None: cfg = self._config(prompt="History : {{hist}} . Profile : {{prof}}") cfg.sid_space.codebook.extend([4, 4, 4]) - compiled = self._compile( - cfg, [create_prompt_feature(_HIST), create_prompt_feature(_PROF)] - ) + compiled = self._compile(cfg, [_feature(_HIST), _feature(_PROF)]) by_name = { s.name: s for s in compiled.prompt_plan.segments if isinstance(s, SlotSeg) @@ -100,7 +100,7 @@ def test_inline_needs_no_group_projected_gets_one(self) -> None: def test_static_runs_are_woven_between_slots(self) -> None: cfg = self._config(prompt="History : {{hist}} . Predict :") cfg.sid_space.codebook.extend([4]) - compiled = self._compile(cfg, [create_prompt_feature(_HIST)]) + compiled = self._compile(cfg, [_feature(_HIST)]) kinds = [ "static" if isinstance(s, Static) else s.name for s in compiled.prompt_plan.segments @@ -112,7 +112,7 @@ def test_static_runs_are_woven_between_slots(self) -> None: def test_scalar_slot_is_one_deep_position(self) -> None: cfg = self._config(prompt="Profile : {{age}}") cfg.sid_space.codebook.extend([4]) - compiled = self._compile(cfg, [create_prompt_feature(_AGE)]) + compiled = self._compile(cfg, [_feature(_AGE)]) seg = next(s for s in compiled.prompt_plan.segments if isinstance(s, SlotSeg)) self.assertIs(seg.fill, FillMode.PROJECTED) self.assertEqual(seg.output_key, "") @@ -127,7 +127,7 @@ def test_manifest_mismatch_is_fatal(self) -> None: cfg.sid_space.codebook.extend([4, 4, 4]) cfg.sid_space.manifest_path = manifest with self.assertRaisesRegex(ValueError, "does not match the manifest"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_manifest_match_compiles(self) -> None: manifest = os.path.join(self.test_dir, "manifest.json") @@ -137,7 +137,7 @@ def test_manifest_match_compiles(self) -> None: cfg.sid_space.codebook.extend([4, 4, 4]) cfg.sid_space.manifest_path = manifest self.assertEqual( - self._compile(cfg, [create_prompt_feature(_HIST)]).sid_space.num_levels, + self._compile(cfg, [_feature(_HIST)]).sid_space.num_levels, 3, ) @@ -147,9 +147,7 @@ def test_rejects_a_mixed_kind_slot(self) -> None: slot = cfg.slots.add(name="both") slot.feature_names.extend(["hist", "age"]) with self.assertRaisesRegex(ValueError, "mixes sequence and scalar"): - self._compile( - cfg, [create_prompt_feature(_HIST), create_prompt_feature(_AGE)] - ) + self._compile(cfg, [_feature(_HIST), _feature(_AGE)]) def test_rejects_unknown_feature_and_unreferenced_slot(self) -> None: cfg = self._config(prompt="X : {{hist}}") @@ -157,13 +155,13 @@ def test_rejects_unknown_feature_and_unreferenced_slot(self) -> None: slot = cfg.slots.add(name="hist") slot.feature_names.append("nope") with self.assertRaisesRegex(ValueError, "not in\n?\\s*feature_configs"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) cfg2 = self._config(prompt="X : {{hist}}") cfg2.sid_space.codebook.extend([4]) cfg2.slots.add(name="ghost").feature_names.append("hist") with self.assertRaisesRegex(ValueError, "never referenced"): - self._compile(cfg2, [create_prompt_feature(_HIST)]) + self._compile(cfg2, [_feature(_HIST)]) def test_rejects_a_projection_on_an_inline_slot(self) -> None: cfg = self._config(prompt="X : {{hist}}") @@ -172,7 +170,7 @@ def test_rejects_a_projection_on_an_inline_slot(self) -> None: slot.feature_names.append("hist") slot.projection.bias = True with self.assertRaisesRegex(ValueError, "is INLINE"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_sid_tokens_absent_from_the_base_tokenizer(self) -> None: cfg = self._config(prompt="X : {{hist}}") @@ -180,14 +178,14 @@ def test_sid_tokens_absent_from_the_base_tokenizer(self) -> None: # renders Histor0..Histor3, and Histor0 is already in the base vocab cfg.sid_space.token_format = "Histor{i}" with self.assertRaisesRegex(ValueError, "already in the base tokenizer"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_a_training_compile_persists_nothing(self) -> None: cfg = self._config(prompt="History : {{hist}}") cfg.sid_space.codebook.extend([4, 4]) before = sorted(os.listdir(self.test_dir)) - compile_prompt(cfg, [create_prompt_feature(_HIST)], ["answer"]) + compile_prompt(cfg, [_feature(_HIST)], ["answer"]) self.assertEqual(sorted(os.listdir(self.test_dir)), before) @@ -195,9 +193,7 @@ def test_extended_tokenizer_is_written(self) -> None: cfg = self._config(prompt="History : {{hist}}") cfg.sid_space.codebook.extend([4, 4]) out = os.path.join(self.test_dir, "export") - compile_prompt( - cfg, [create_prompt_feature(_HIST)], ["answer"], tokenizer_dir=out - ) + compile_prompt(cfg, [_feature(_HIST)], ["answer"], tokenizer_dir=out) written = os.path.join(out, "tokenizer.json") self.assertTrue(os.path.exists(written)) # the SID tokens round-trip, which is what serving reloads @@ -208,7 +204,7 @@ def test_extended_tokenizer_is_written(self) -> None: def test_answer_width_comes_from_the_codebook(self) -> None: cfg = self._config(prompt="History : {{hist}}", response="{{answer}}") cfg.sid_space.codebook.extend([4, 4, 4]) - compiled = self._compile(cfg, [create_prompt_feature(_HIST)]) + compiled = self._compile(cfg, [_feature(_HIST)]) seg = next( s for s in compiled.prompt_plan.response_segments if isinstance(s, SlotSeg) @@ -229,9 +225,7 @@ def test_response_must_name_a_label_field(self) -> None: r"\[prof\] names \['prof'\], which are not in " r"data_config.label_fields", ): - self._compile( - cfg, [create_prompt_feature(_HIST), create_prompt_feature(_PROF)] - ) + self._compile(cfg, [_feature(_HIST), _feature(_PROF)]) def test_response_slot_takes_exactly_one_label_field(self) -> None: cfg = self._config(prompt="History : {{hist}}", response="{{answer}}") @@ -240,7 +234,7 @@ def test_response_slot_takes_exactly_one_label_field(self) -> None: slot.feature_names.extend(["sid_a", "sid_b"]) with self.assertRaisesRegex(ValueError, "is one label field"): - compile_prompt(cfg, [create_prompt_feature(_HIST)], ["sid_a", "sid_b"]) + compile_prompt(cfg, [_feature(_HIST)], ["sid_a", "sid_b"]) def test_response_slot_may_not_declare_a_projection(self) -> None: cfg = self._config(prompt="History : {{hist}}", response="{{answer}}") @@ -250,14 +244,14 @@ def test_response_slot_may_not_declare_a_projection(self) -> None: slot.projection.SetInParent() with self.assertRaisesRegex(ValueError, "drop its projection"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_missing_sid_space_is_rejected(self) -> None: # the response width is codebook-derived, so sid_space must exist cfg = self._config(prompt="History : {{hist}}", response="{{answer}}") cfg.ClearField("sid_space") with self.assertRaisesRegex(ValueError, "sid_space is required"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_token_format_without_a_placeholder_is_rejected(self) -> None: # without {i} every token renders alike: one row, not sum(codebook) @@ -265,13 +259,13 @@ def test_token_format_without_a_placeholder_is_rejected(self) -> None: cfg.sid_space.codebook.extend([4, 4, 4]) cfg.sid_space.token_format = "<|sid|>" with self.assertRaisesRegex(ValueError, "has no '{i}' placeholder"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_a_custom_token_format_with_a_placeholder_compiles(self) -> None: cfg = self._config(prompt="History : {{hist}}", response="{{answer}}") cfg.sid_space.codebook.extend([4, 4, 4]) cfg.sid_space.token_format = "C{i}" - compiled = self._compile(cfg, [create_prompt_feature(_HIST)]) + compiled = self._compile(cfg, [_feature(_HIST)]) space = compiled.sid_space self.assertEqual(space.band_hi[-1] - space.band_lo[0] + 1, 12) @@ -282,7 +276,7 @@ def test_missing_response_is_rejected(self) -> None: cfg.sid_space.codebook.extend([4, 4, 4]) cfg.ClearField("response") with self.assertRaisesRegex(ValueError, "response is required"): - self._compile(cfg, [create_prompt_feature(_HIST)]) + self._compile(cfg, [_feature(_HIST)]) def test_a_grouped_feature_inherits_the_group_cap(self) -> None: # a SequenceFeature member never sets its own sequence_length; the cap diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 5c103d43..0f7febf1 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -32,14 +32,9 @@ from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.tests import utils -from tzrec.tests.prompt_test_util import ( - _CODEBOOK, - _WORDS, - create_prompt_tokenizer, - projected_feature, -) from tzrec.utils import config_util from tzrec.utils.test_util import ( + create_genrec_test_tokenizer, create_tiny_causal_lm, gpu_unavailable, make_test_dir, @@ -48,6 +43,11 @@ _MOCK_CONFIG = "tzrec/tests/configs/genrec_causal_lm_model_mock.config" _BUNDLE_UUID = "bundle-test" +_CODEBOOK = [4, 4, 4] +_BEH = ( + 'sequence_id_feature { feature_name: "beh" expression: "user:beh" ' + "num_buckets: 32 embedding_dim: 8 sequence_length: 2 }" +) class GenRecIntegrationTest(unittest.TestCase): @@ -69,8 +69,8 @@ def _prepare_config(self, projected: bool) -> str: """Write the tiny backbone, tokenizer, manifest and data; return the config.""" backbone = os.path.join(self.test_dir, "backbone") create_tiny_causal_lm(64).save_pretrained(backbone) - tokenizer = create_prompt_tokenizer( - os.path.join(self.test_dir, "tok.json"), _WORDS + tokenizer = create_genrec_test_tokenizer( + os.path.join(self.test_dir, "tok.json") ) manifest = os.path.join(self.test_dir, "manifest.json") with open(manifest, "w") as f: @@ -86,7 +86,7 @@ def _prepare_config(self, projected: bool) -> str: config.prompt_config.tokenizer_path = tokenizer config.prompt_config.sid_space.manifest_path = manifest if projected: - text_format.Merge(projected_feature("beh", 8), config.feature_configs.add()) + text_format.Merge(_BEH, config.feature_configs.add()) config.prompt_config.prompt = "History : {{hist}} . {{beh}} Predict :" config_path = os.path.join(self.test_dir, "genrec.config") config_util.save_message(config, config_path) diff --git a/tzrec/tests/prompt_test_util.py b/tzrec/tests/prompt_test_util.py deleted file mode 100644 index 02e35cdb..00000000 --- a/tzrec/tests/prompt_test_util.py +++ /dev/null @@ -1,163 +0,0 @@ -# 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 unittest -from typing import Any, Dict, Sequence - -import numpy as np -import torch -from google.protobuf import text_format -from tokenizers import Tokenizer, models, pre_tokenizers - -from tzrec.datasets.utils import BASE_DATA_GROUP, Batch -from tzrec.features.feature import BaseFeature, FgMode, create_features -from tzrec.main import _create_model -from tzrec.prompt.assembler import PromptAssembler -from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.types import CompiledPrompt -from tzrec.protos import feature_pb2 -from tzrec.protos.model_pb2 import ModelConfig -from tzrec.protos.prompt_pb2 import PromptConfig -from tzrec.utils.test_util import create_tiny_causal_lm, make_test_dir - - -def create_prompt_tokenizer(path: str, words: Sequence[str]) -> str: - """Write a word-level tokenizer used by prompt tests. - - Args: - path: destination JSON path. - words: vocabulary entries in token-id order. - - Returns: - The destination path. - """ - tokenizer = Tokenizer( - models.WordLevel( - vocab={word: i for i, word in enumerate(words)}, unk_token="" - ) - ) - # pyrefly: ignore[read-only] - tokenizer.pre_tokenizer = pre_tokenizers.Whitespace() - tokenizer.save(path) - return path - - -def create_prompt_feature(text: str) -> BaseFeature: - """Create one prompt test feature from text-format protobuf. - - Args: - text: text-format ``FeatureConfig``. - - Returns: - The created feature. - """ - config = feature_pb2.FeatureConfig() - text_format.Merge(text, config) - return create_features([config], fg_mode=FgMode.FG_NONE)[0] - - -def offset_sid_codes(codes: Sequence[Any], codebook: Sequence[int]) -> np.ndarray: - """Shift local SID codes into the flattened per-level space. - - Args: - codes: local codes grouped by SID item. - codebook: vocabulary size for each SID level. - - Returns: - Flat offset codes in item-major order. - """ - offsets = np.cumsum([0, *codebook[:-1]]) - return (np.asarray(codes).reshape(-1, len(codebook)) + offsets).reshape(-1) - - -def assemble_into( - compiled_prompt: CompiledPrompt, parsed_features: Dict[str, Any] -) -> Dict[str, torch.Tensor]: - """Assemble one parsed batch the way the collator does. - - Args: - compiled_prompt: the compiled prompt. - parsed_features: ``{column}.values`` / ``{column}.lengths`` as the data - parser emits them, as tensors or arrays. - - Returns: - The assembled streams keyed for ``additional_infos``. - """ - batch = {k: torch.as_tensor(np.asarray(v)) for k, v in parsed_features.items()} - return PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( - batch - ) - - -_CODEBOOK = [4, 4, 4] -_WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] -_HIST = 'sequence_raw_feature { feature_name: "hist" expression: "user:hist" }' - - -def projected_feature(name: str, dim: int) -> str: - """A PROJECTED slot member: a sequence id feature with an embedding.""" - return ( - f'sequence_id_feature {{ feature_name: "{name}" expression: "user:{name}" ' - f"num_buckets: 32 embedding_dim: {dim} sequence_length: 2 }}" - ) - - -class GenRecModelTestBase(unittest.TestCase): - """Builds a real GenRecCausalLMModel over a tiny backbone and prompt.""" - - def setUp(self) -> None: - """Build the tiny backbone, tokenizer and compiled prompt.""" - self.test_dir = make_test_dir() - self.backbone = os.path.join(self.test_dir, "backbone") - create_tiny_causal_lm(64).save_pretrained(self.backbone) - self.tok = create_prompt_tokenizer( - os.path.join(self.test_dir, "tok.json"), _WORDS - ) - self.features = [create_prompt_feature(_HIST)] - self.compiled_prompt = self._compile(self.features) - - def _compile(self, features, template="History : {{hist}} . Predict :", **kwargs): - kwargs.setdefault("response", "{{answer}}") - cfg = PromptConfig(tokenizer_path=self.tok, prompt=template, **kwargs) - cfg.sid_space.codebook.extend(_CODEBOOK) - return compile_prompt(cfg, features, ["answer"]) - - def _model( - self, - features=None, - compiled_prompt=-1, - beam_widths=(2, 2, 2), - num_return_sequences=2, - lm_parameter_dtype=None, - hf_model_name_or_path=None, - ): - model_config = ModelConfig() - lm_cfg = model_config.genrec_causal_lm_model - lm_cfg.hf_model_name_or_path = hf_model_name_or_path or self.backbone - lm_cfg.common.beam_widths.extend(beam_widths) - lm_cfg.common.num_return_sequences = num_return_sequences - if lm_parameter_dtype is not None: - lm_cfg.common.lm_parameter_dtype = lm_parameter_dtype - return _create_model( - model_config, - self.features if features is None else features, - ["answer"], - compiled_prompt=( - self.compiled_prompt if compiled_prompt == -1 else compiled_prompt - ), - ) - - def _batch(self, parsed, compiled_prompt=None, sparse=None): - streams = assemble_into(compiled_prompt or self.compiled_prompt, parsed) - batch = Batch(sparse_features={BASE_DATA_GROUP: sparse} if sparse else {}) - batch.additional_infos.update(streams) - return batch diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 440f555d..7d221e5b 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -14,7 +14,7 @@ import os import tempfile from enum import Enum -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Dict, List, Optional, Sequence, Tuple, Union import numpy as np import pandas as pd @@ -25,7 +25,11 @@ from torch.fx import GraphModule from tzrec.acc.aot_utils import export_model_aot, load_model_aot -from tzrec.models.model import ScriptWrapper +from tzrec.models.model import BaseModel, ScriptWrapper +from tzrec.prompt.types import CompiledPrompt +from tzrec.protos import feature_pb2 +from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig +from tzrec.protos.prompt_pb2 import PromptSlot from tzrec.utils.export_util import split_model from tzrec.utils.fx_util import symbolic_trace @@ -350,3 +354,98 @@ def reference_stu_truncation( chunks.append(torch.cat([prefix, uih_kept, targets], dim=0)) new_lens.append(contextual_seq_len + new_uih + T) return torch.cat(chunks, dim=0), new_lens + + +def create_genrec_test_tokenizer( + path: str, + words: Sequence[str] = ("History", "Predict", ":", ".", "", "<|im_end|>"), +) -> str: + """Write a word-level tokenizer for the genrec tests. + + Args: + path (str): destination JSON path. + words (Sequence[str]): vocabulary entries in token-id order. + + Returns: + str: the destination path. + """ + from tokenizers import Tokenizer, models, pre_tokenizers + + tokenizer = Tokenizer( + models.WordLevel( + vocab={word: i for i, word in enumerate(words)}, unk_token="" + ) + ) + # pyrefly: ignore[read-only] + tokenizer.pre_tokenizer = pre_tokenizers.Whitespace() + tokenizer.save(path) + return path + + +def create_genrec_test_model( + test_dir: str, + feature_configs: Optional[List[feature_pb2.FeatureConfig]] = None, + prompt: str = "History : {{hist}} . Predict :", + response: str = "{{answer}}", + slots: Sequence[PromptSlot] = (), + beam_widths: Sequence[int] = (2, 2, 2), + num_return_sequences: int = 2, + lm_parameter_dtype: Optional["GenRecModelConfig.ParamDtype"] = None, +) -> Tuple[BaseModel, CompiledPrompt]: + """Build a GenRecCausalLMModel over a tiny backbone and a compiled prompt. + + The backbone and the tokenizer are written under ``test_dir``. The default + features are one ``hist`` raw sequence and the codebook is ``(4, 4, 4)``, so + an offset SID code is ``level_offsets[l] + code`` with offsets ``(0, 4, 8)``. + + Args: + test_dir (str): scratch directory. + feature_configs (list, optional): feature configs; the ``hist`` raw + sequence when None. + prompt (str): the prompt template. + response (str): the response template. + slots (Sequence[PromptSlot]): explicit slot declarations. + beam_widths (Sequence[int]): per-level beam widths. + num_return_sequences (int): sequences returned per sample. + lm_parameter_dtype (optional): ``GenRecModelConfig.ParamDtype`` value. + + Returns: + Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was built on. + """ + from tzrec.features.feature import FgMode, create_features + from tzrec.main import _create_model + from tzrec.prompt.compile import compile_prompt + from tzrec.protos.model_pb2 import ModelConfig + from tzrec.protos.prompt_pb2 import PromptConfig + + backbone = os.path.join(test_dir, "backbone") + create_tiny_causal_lm(64).save_pretrained(backbone) + if feature_configs is None: + feature_configs = [ + feature_pb2.FeatureConfig( + sequence_raw_feature=feature_pb2.RawFeature( + feature_name="hist", expression="user:hist" + ) + ) + ] + features = create_features(feature_configs, fg_mode=FgMode.FG_NONE) + prompt_config = PromptConfig( + tokenizer_path=create_genrec_test_tokenizer(os.path.join(test_dir, "tok.json")), + prompt=prompt, + response=response, + ) + prompt_config.sid_space.codebook.extend([4, 4, 4]) + prompt_config.slots.extend(slots) + compiled_prompt = compile_prompt(prompt_config, features, ["answer"]) + + model_config = ModelConfig() + lm_config = model_config.genrec_causal_lm_model + lm_config.hf_model_name_or_path = backbone + lm_config.common.beam_widths.extend(beam_widths) + lm_config.common.num_return_sequences = num_return_sequences + if lm_parameter_dtype is not None: + lm_config.common.lm_parameter_dtype = lm_parameter_dtype + model = _create_model( + model_config, features, ["answer"], compiled_prompt=compiled_prompt + ) + return model, compiled_prompt From c88e2dafbdbdba8b95c23f8a62de9ebff90f8166 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 15:49:24 +0800 Subject: [PATCH 14/21] [refactor] write one config.json for the genrec export The export composes config.json once from the backbone config dcp_to_hf returns, under the names GenRecForCausalLM / genrec, and writes the tokenizer beside it. The server never reads the SID space and the offline index builder derives it from the tokenizer and the bundle manifest, so prompt/prompt.json and persist.py are gone. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/prompt/persist.py | 72 --------------------- tzrec/tests/genrec_integration_test.py | 36 +++++------ tzrec/utils/hf_export_util.py | 88 +++++++++----------------- tzrec/utils/hf_export_util_test.py | 6 +- 4 files changed, 52 insertions(+), 150 deletions(-) delete mode 100644 tzrec/prompt/persist.py diff --git a/tzrec/prompt/persist.py b/tzrec/prompt/persist.py deleted file mode 100644 index c99a7414..00000000 --- a/tzrec/prompt/persist.py +++ /dev/null @@ -1,72 +0,0 @@ -# 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. - -"""Publishes the SID space beside an export. - -Export writes ``prompt/prompt.json`` with the resolved SID space: the token -base, the per-level bands and the bundle identity a serving side needs to -build its constraint index and to refuse one built from another bundle. The -plan itself is deliberately not published -- it reaches serving compiled into -the front-end, and a copy a runtime could interpret would invite the second -assembler this design exists to prevent. -""" - -import dataclasses -import json -import os - -from tzrec.prompt.types import ResolvedSidSpace - -PROMPT_DIR = "prompt" -PROMPT_CONTRACT_FILENAME = "prompt.json" -TOKENIZER_DIR = "tokenizer" - - -def read_bundle_uuid(manifest_path: str) -> str: - """The identity of the SID bundle a manifest describes, empty without one. - - Args: - manifest_path: ``sid_space.manifest_path``, possibly unset. - - Returns: - The manifest's ``bundle_uuid``, or ``""`` when there is no manifest or - it records none. - """ - if not manifest_path: - return "" - with open(manifest_path, "r") as f: - return str(json.load(f).get("bundle_uuid", "")) - - -def write_serving_contract( - sid_space: ResolvedSidSpace, bundle_uuid: str, export_dir: str -) -> str: - """Write ``prompt/prompt.json``: the resolved SID space and the bundle identity. - - Args: - sid_space: the resolved SID token space. - bundle_uuid: the bundle the codebook was read from, so an index builder - can refuse a catalog from another bundle; empty when unknown. - export_dir: the export directory. - - Returns: - The path written. - """ - out = os.path.join(export_dir, PROMPT_DIR) - os.makedirs(out, exist_ok=True) - path = os.path.join(out, PROMPT_CONTRACT_FILENAME) - with open(path, "w") as f: - json.dump( - {"sid_space": dataclasses.asdict(sid_space), "bundle_uuid": bundle_uuid}, - f, - indent=2, - ) - return path diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 0f7febf1..a69b25c1 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -131,9 +131,8 @@ def test_genrec_train_eval_export(self): "model_acc.json", "config.json", "model.safetensors", - "prompt/prompt.json", - "prompt/tokenizer/tokenizer.json", - "prompt/tokenizer/tokenizer_config.json", + "tokenizer.json", + "tokenizer_config.json", ): self.assertTrue(os.path.exists(os.path.join(export_dir, name)), name) self.assertFalse(os.path.exists(os.path.join(export_dir, "sparse"))) @@ -171,26 +170,23 @@ def test_genrec_train_eval_export(self): # what the engine loads beside it with open(os.path.join(export_dir, "config.json"), "r") as f: hf_config = json.load(f) - self.assertEqual(hf_config["architectures"], ["PromptGenRecForCausalLM"]) - self.assertEqual(hf_config["model_type"], "prompt_genrec") - self.assertEqual(hf_config["text_config"]["model_type"], "qwen2") - self.assertEqual( - hf_config["text_config"]["vocab_size"], compiled.sid_space.target_vocab_size - ) - with open(os.path.join(export_dir, "prompt", "prompt.json"), "r") as f: - contract = json.load(f) - self.assertEqual(list(contract), ["sid_space", "bundle_uuid"]) - self.assertEqual(contract["bundle_uuid"], _BUNDLE_UUID) + self.assertEqual(hf_config["architectures"], ["GenRecForCausalLM"]) + self.assertEqual(hf_config["model_type"], "genrec") self.assertEqual( - contract["sid_space"]["band_lo"], list(compiled.sid_space.band_lo) - ) - self.assertEqual( - contract["sid_space"]["band_hi"], list(compiled.sid_space.band_hi) + hf_config["text_config"]["architectures"], ["Qwen2ForCausalLM"] ) + self.assertEqual(hf_config["text_config"]["model_type"], "qwen2") + self.assertEqual(hf_config["vocab_size"], compiled.sid_space.target_vocab_size) + self.assertEqual(hf_config["eos_token_id"], compiled.sid_space.eos_token_id) + self.assertEqual(hf_config["pad_token_id"], compiled.sid_space.pad_token_id) + self.assertNotIn("sid_space", hf_config) from transformers import AutoTokenizer - tokenizer = AutoTokenizer.from_pretrained( - os.path.join(export_dir, "prompt", "tokenizer") + # the index builder derives the token base from the first SID token + tokenizer = AutoTokenizer.from_pretrained(export_dir) + self.assertEqual( + tokenizer.convert_tokens_to_ids("<|sid_0|>"), + compiled.sid_space.base_vocab_size, ) self.assertEqual(tokenizer.decode([compiled.sid_space.band_lo[0]]), "<|sid_0|>") self.assertEqual( @@ -218,7 +214,7 @@ def test_genrec_export_distributed_embedding(self): "sparse_features.json", ): self.assertTrue(os.path.exists(os.path.join(sparse_dir, name)), name) - for name in ("config.json", "model.safetensors", "prompt/prompt.json"): + for name in ("config.json", "model.safetensors", "tokenizer.json"): self.assertTrue(os.path.exists(os.path.join(dist_dir, name)), name) with open(os.path.join(dist_dir, "model_acc.json"), "r") as f: self.assertEqual(json.load(f)["DISTRIBUTED_EMBEDDING"], "1") diff --git a/tzrec/utils/hf_export_util.py b/tzrec/utils/hf_export_util.py index f67c32dc..56ac98c8 100644 --- a/tzrec/utils/hf_export_util.py +++ b/tzrec/utils/hf_export_util.py @@ -28,22 +28,15 @@ from tzrec.constant import HF_EXPORT_META_FILENAME from tzrec.features.feature import BaseFeature from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.persist import ( - PROMPT_DIR, - TOKENIZER_DIR, - read_bundle_uuid, - write_serving_contract, -) from tzrec.protos.pipeline_pb2 import EasyRecConfig from tzrec.utils import checkpoint_util from tzrec.utils.filesystem_util import url_to_fs from tzrec.utils.logging_util import logger -SERVING_ARCH = "PromptGenRecForCausalLM" -SERVING_MODEL_TYPE = "prompt_genrec" +SERVING_ARCH = "GenRecForCausalLM" +SERVING_MODEL_TYPE = "genrec" _HF_ASSET_FILES = ( - "config.json", "generation_config.json", "tokenizer.json", "tokenizer_config.json", @@ -86,44 +79,15 @@ def write_hf_assets(wrapped_model: nn.Module, save_dir: str) -> None: json.dump(meta, f, indent=2) -def write_composite_config(export_dir: str) -> None: - """Rewrite ``config.json`` so the backbone sits under ``text_config``. - - That is what names the backbone to a serving runtime that composes an - arbitrary causal LM behind one registered architecture, and what makes it - treat the model as carrying a second modality, which is what a projected - prompt slot is. - - Args: - export_dir: the HuggingFace export directory. - """ - path = os.path.join(export_dir, "config.json") - with open(path, "r") as f: - backbone: Dict[str, Any] = json.load(f) - if backbone.get("model_type") == SERVING_MODEL_TYPE: - return - composite: Dict[str, Any] = { - "architectures": [SERVING_ARCH], - "model_type": SERVING_MODEL_TYPE, - "text_config": backbone, - } - # a runtime that reads only the outer config still needs to size its cache - for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): - if key in backbone: - composite[key] = backbone[key] - with open(path, "w") as f: - json.dump(composite, f, indent=2) - logger.info( - f"wrote a composite config naming backbone " - f"{backbone.get('architectures', ['?'])[0]} under {SERVING_ARCH}." - ) - - -def dcp_to_hf(ckpt_dir: str, out_dir: str) -> None: +def dcp_to_hf(ckpt_dir: str, out_dir: str) -> Dict[str, Any]: """Convert a checkpoint with co-located HF assets to a ``from_pretrained`` dir. Keys that do not map 1:1 onto the co-located ``config.json`` raise rather - than write a partial model. + than write a partial model. The config itself is not written: the caller + composes the one ``config.json`` the export carries. + + Returns: + The backbone config as ``config.json`` would hold it. """ from torch.distributed.checkpoint.state_dict_loader import ( _load_state_dict_from_keys, @@ -203,6 +167,7 @@ def _derive_by_suffix( src = os.path.join(ckpt_dir, fname) if os.path.exists(src): shutil.copy(src, os.path.join(out_dir, fname)) + return json.loads(cfg.to_json_string()) def export_hf_assets( @@ -213,9 +178,12 @@ def export_hf_assets( ) -> None: """Write what an LLM engine reads beside the scripted front-end. - The HuggingFace weights and composite config, the extended tokenizer under - ``prompt/tokenizer`` and ``prompt/prompt.json``. A remote ``export_dir`` is - written locally and uploaded, as ``export_model`` does for its own files. + The HuggingFace weights, the extended tokenizer and one composite + ``config.json``: the backbone sits under ``text_config``, which names it to + a runtime that composes an arbitrary causal LM behind one registered + architecture and treats a projected prompt slot as a second modality. A + remote ``export_dir`` is written locally and uploaded, as ``export_model`` + does for its own files. Args: pipeline_config: the pipeline being exported. @@ -226,20 +194,26 @@ def export_hf_assets( fs, local_dir = url_to_fs(export_dir) if fs is not None: local_dir = tempfile.mkdtemp() - dcp_to_hf(checkpoint_path, local_dir) - write_composite_config(local_dir) - prompt_config = pipeline_config.prompt_config + backbone = dcp_to_hf(checkpoint_path, local_dir) compiled = compile_prompt( - prompt_config, + pipeline_config.prompt_config, features, list(pipeline_config.data_config.label_fields), - tokenizer_dir=os.path.join(local_dir, PROMPT_DIR, TOKENIZER_DIR), - ) - write_serving_contract( - compiled.sid_space, - read_bundle_uuid(prompt_config.sid_space.manifest_path), - local_dir, + tokenizer_dir=local_dir, ) + composite: Dict[str, Any] = { + "architectures": [SERVING_ARCH], + "model_type": SERVING_MODEL_TYPE, + "text_config": backbone, + "eos_token_id": compiled.sid_space.eos_token_id, + "pad_token_id": compiled.sid_space.pad_token_id, + } + # a runtime that reads only the outer config still needs to size its cache + for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): + if key in backbone: + composite[key] = backbone[key] + with open(os.path.join(local_dir, "config.json"), "w") as f: + json.dump(composite, f, indent=2) if fs is not None: fs.upload(local_dir, export_dir, recursive=True, file_thread_num=os.cpu_count()) shutil.rmtree(local_dir) diff --git a/tzrec/utils/hf_export_util_test.py b/tzrec/utils/hf_export_util_test.py index 354e73d3..bff59edf 100644 --- a/tzrec/utils/hf_export_util_test.py +++ b/tzrec/utils/hf_export_util_test.py @@ -131,7 +131,11 @@ def test_dcp_to_hf_round_trip_drops_tied_head(self) -> None: lm = _tied_lm() ckpt_dir = self._save_ckpt(_DmpLike(_TrainWrapper(_GenRec(lm)))) out_dir = os.path.join(self.test_dir, "hf_out") - dcp_to_hf(ckpt_dir, out_dir) + config = dcp_to_hf(ckpt_dir, out_dir) + # the caller composes config.json; here the backbone's own is enough + self.assertEqual(config["model_type"], lm.config.model_type) + with open(os.path.join(out_dir, "config.json"), "w") as f: + json.dump(config, f) st = load_file(os.path.join(out_dir, "model.safetensors")) self.assertNotIn("lm_head.weight", st) From 5a169f09219685f3d3a2910551b47f00a625cae0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 15:49:51 +0800 Subject: [PATCH 15/21] [chore] bump version to 1.4.5 Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBqWYDaN --- tzrec/version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tzrec/version.py b/tzrec/version.py index c460f30b..84ae177a 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.4.3" +__version__ = "1.4.5" From 48adf079314a9b18f0edb6718ad00f67b722b1c0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:05:07 +0800 Subject: [PATCH 16/21] [perf] load only the backbone at genrec export The export built the LM on every rank and kept it through the shard, and dcp_to_hf read every checkpoint key -- all the sparse tables -- to keep the backbone's. The front-end holds no LM, so it is freed once built, and the DCP key mapping now runs on the metadata's names so the load asks for the backbone alone. The staging dir moves into a try/finally. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/main.py | 6 +- tzrec/utils/hf_export_util.py | 95 ++++++++++++++++-------------- tzrec/utils/hf_export_util_test.py | 18 ++++++ 3 files changed, 75 insertions(+), 44 deletions(-) diff --git a/tzrec/main.py b/tzrec/main.py index c4c83cd2..661a476d 100644 --- a/tzrec/main.py +++ b/tzrec/main.py @@ -1278,9 +1278,13 @@ def export( ) if not checkpoint_path: raise ValueError("HF export: no checkpoint found to convert.") + front_end = InferWrapper(GenRecFrontEnd(model.model)) + # the front-end carries no backbone: the engine serves the LM from the + # HuggingFace weights, so the copy built here is dead weight on every rank + del model.model.lm export_model( ori_pipeline_config, - InferWrapper(GenRecFrontEnd(model.model)), + front_end, checkpoint_path, export_dir, assets=assets, diff --git a/tzrec/utils/hf_export_util.py b/tzrec/utils/hf_export_util.py index 56ac98c8..0597862a 100644 --- a/tzrec/utils/hf_export_util.py +++ b/tzrec/utils/hf_export_util.py @@ -91,6 +91,7 @@ def dcp_to_hf(ckpt_dir: str, out_dir: str) -> Dict[str, Any]: """ from torch.distributed.checkpoint.state_dict_loader import ( _load_state_dict_from_keys, + _storage_setup, ) from transformers import AutoConfig, AutoModelForCausalLM @@ -98,11 +99,6 @@ def dcp_to_hf(ckpt_dir: str, out_dir: str) -> Dict[str, Any]: if not os.path.exists(model_ckpt_path): raise RuntimeError(f"dcp_to_hf: model DCP dir [{model_ckpt_path}] not exists.") - # No keys => load every key; non-distributed => full tensors locally. - raw_state: Dict[str, torch.Tensor] = _load_state_dict_from_keys( - checkpoint_id=model_ckpt_path - ) - meta_path = os.path.join(ckpt_dir, HF_EXPORT_META_FILENAME) prefix: Optional[str] = None if os.path.exists(meta_path): @@ -116,45 +112,52 @@ def dcp_to_hf(ckpt_dir: str, out_dir: str) -> Dict[str, Any]: tied_keys: Set[str] = set(getattr(empty, "_tied_weights_keys", None) or []) del empty - def _strip_recorded_prefix( - state: Dict[str, torch.Tensor], - ) -> Optional[Dict[str, torch.Tensor]]: + # the mapping is decided on names alone, so the load below reads the + # backbone and not the sparse tables beside it + reader = _storage_setup(None, model_ckpt_path, reader=True) + ckpt_keys: Set[str] = set(reader.read_metadata().state_dict_metadata) + + def _strip_recorded_prefix() -> Optional[Dict[str, str]]: """Strip the recorded prefix; None unless it yields an EXACT match.""" if not prefix: return None - out = {k[len(prefix) :]: v for k, v in state.items() if k.startswith(prefix)} + out = {k[len(prefix) :]: k for k in ckpt_keys if k.startswith(prefix)} return out if set(out) == target_keys else None - def _derive_by_suffix( - state: Dict[str, torch.Tensor], - ) -> Optional[Dict[str, torch.Tensor]]: + def _derive_by_suffix() -> Optional[Dict[str, str]]: """Each target key is a unique suffix of exactly one DCP key; None if not.""" - out: Dict[str, torch.Tensor] = {} + out: Dict[str, str] = {} for tk in target_keys: - matches = [k for k in state if k == tk or k.endswith("." + tk)] + matches = [k for k in ckpt_keys if k == tk or k.endswith("." + tk)] if len(matches) != 1: return None - out[tk] = state[matches[0]] + out[tk] = matches[0] return out - mapped = _strip_recorded_prefix(raw_state) - if mapped is None: + key_map = _strip_recorded_prefix() + if key_map is None: if prefix: logger.warning( f"dcp_to_hf: recorded prefix [{prefix}] did not map exactly onto " "the architecture; deriving the backbone prefix by suffix-matching." ) - mapped = _derive_by_suffix(raw_state) + key_map = _derive_by_suffix() - if mapped is None: + if key_map is None: raise RuntimeError( "dcp_to_hf: cannot map the DCP state dict onto the backbone " f"architecture (recorded prefix={prefix!r}). Wanted " f"{len(target_keys)} keys like {sorted(target_keys)[:3]}; the " - f"checkpoint holds {len(raw_state)} like {sorted(raw_state)[:3]}. " + f"checkpoint holds {len(ckpt_keys)} like {sorted(ckpt_keys)[:3]}. " "Refusing to write a partially-loaded HF model." ) + # non-distributed => full tensors locally + raw_state: Dict[str, torch.Tensor] = _load_state_dict_from_keys( + set(key_map.values()), checkpoint_id=model_ckpt_path + ) + mapped = {tk: raw_state[ck] for tk, ck in key_map.items()} + # from_pretrained re-ties them. if getattr(cfg, "tie_word_embeddings", False): mapped = {k: v for k, v in mapped.items() if k not in tied_keys} @@ -194,26 +197,32 @@ def export_hf_assets( fs, local_dir = url_to_fs(export_dir) if fs is not None: local_dir = tempfile.mkdtemp() - backbone = dcp_to_hf(checkpoint_path, local_dir) - compiled = compile_prompt( - pipeline_config.prompt_config, - features, - list(pipeline_config.data_config.label_fields), - tokenizer_dir=local_dir, - ) - composite: Dict[str, Any] = { - "architectures": [SERVING_ARCH], - "model_type": SERVING_MODEL_TYPE, - "text_config": backbone, - "eos_token_id": compiled.sid_space.eos_token_id, - "pad_token_id": compiled.sid_space.pad_token_id, - } - # a runtime that reads only the outer config still needs to size its cache - for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): - if key in backbone: - composite[key] = backbone[key] - with open(os.path.join(local_dir, "config.json"), "w") as f: - json.dump(composite, f, indent=2) - if fs is not None: - fs.upload(local_dir, export_dir, recursive=True, file_thread_num=os.cpu_count()) - shutil.rmtree(local_dir) + try: + backbone = dcp_to_hf(checkpoint_path, local_dir) + compiled = compile_prompt( + pipeline_config.prompt_config, + features, + list(pipeline_config.data_config.label_fields), + tokenizer_dir=local_dir, + ) + composite: Dict[str, Any] = { + "architectures": [SERVING_ARCH], + "model_type": SERVING_MODEL_TYPE, + "text_config": backbone, + "eos_token_id": compiled.sid_space.eos_token_id, + "pad_token_id": compiled.sid_space.pad_token_id, + } + # a runtime that reads only the outer config still needs to size its cache + for key in ("vocab_size", "hidden_size", "num_hidden_layers", "torch_dtype"): + if key in backbone: + composite[key] = backbone[key] + with open(os.path.join(local_dir, "config.json"), "w") as f: + json.dump(composite, f, indent=2) + if fs is not None: + fs.upload( + local_dir, export_dir, recursive=True, file_thread_num=os.cpu_count() + ) + finally: + # the staging dir holds the full LM weights; drop it however this ends + if fs is not None: + shutil.rmtree(local_dir, ignore_errors=True) diff --git a/tzrec/utils/hf_export_util_test.py b/tzrec/utils/hf_export_util_test.py index bff59edf..585d4998 100644 --- a/tzrec/utils/hf_export_util_test.py +++ b/tzrec/utils/hf_export_util_test.py @@ -147,6 +147,24 @@ def test_dcp_to_hf_round_trip_drops_tied_head(self) -> None: for k, v in lm.state_dict().items(): self.assertTrue(torch.equal(back.state_dict()[k], v), k) + def test_dcp_to_hf_loads_only_the_backbone_keys(self) -> None: + """The rest of a genrec checkpoint is the sparse tables; never read them.""" + from torch.distributed.checkpoint import state_dict_loader + + ckpt_dir = self._save_ckpt(_TrainWrapper(_GenRec(_tied_lm()))) + original = state_dict_loader._load_state_dict_from_keys + requested = [] + + def _spy(keys=None, **kwargs): + requested.append(keys) + return original(keys, **kwargs) + + with mock.patch.object(state_dict_loader, "_load_state_dict_from_keys", _spy): + dcp_to_hf(ckpt_dir, os.path.join(self.test_dir, "hf_out_keys")) + self.assertEqual(len(requested), 1) + self.assertIsNotNone(requested[0]) + self.assertFalse([k for k in requested[0] if ".other." in k]) + def test_dcp_to_hf_refuses_a_mismatched_architecture(self) -> None: ckpt_dir = self._save_ckpt(_TrainWrapper(_GenRec(_tied_lm()))) # widen the recorded architecture so the checkpoint can no longer fill it From a49fda2abdc0b223de0edde4e872d22d052e5a59 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:05:09 +0800 Subject: [PATCH 17/21] [refactor] seed each hole key with its slot salt An empty hole folded to zero in every slot, which the fold's own invariant says cannot happen: the slot must discriminate. The accumulator starts at the salt, and the salt is 1-based so slot 0 is salted too. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/prompt/hole_keys.py | 12 ++++++++++-- tzrec/prompt/hole_keys_test.py | 18 ++++++++++++++++++ 2 files changed, 28 insertions(+), 2 deletions(-) diff --git a/tzrec/prompt/hole_keys.py b/tzrec/prompt/hole_keys.py index bf0942ff..00e06163 100644 --- a/tzrec/prompt/hole_keys.py +++ b/tzrec/prompt/hole_keys.py @@ -68,6 +68,10 @@ class HoleKeyBuilder(nn.Module): and all silent. A dense member contributes its float32 bit pattern per row, which is the parsed input and not a computed reduction. + Body and response slots are walked together, which the assembler's holes + match because compile gives every response slot ``FillMode.INLINE``: a + response slot is never PROJECTED and so never reaches ``projected_slots``. + Args: prompt_plan: the compiled plan; its ``projected_slots`` fix the hole order. @@ -94,7 +98,9 @@ def __init__(self, prompt_plan: PromptPlan) -> None: for seg in prompt_plan.projected_slots: self.member_names.append(list(seg.feature_names)) self.is_sequences.append(seg.group_type == FeatureGroupType.JAGGED_SEQUENCE) - self.salts.append(_wrap64(self.C_SLOT * int(seg.slot_id))) + # 1-based: slot 0 must not salt to zero, or its empty hole + # would fold to zero as an unsalted accumulator did + self.salts.append(_wrap64(self.C_SLOT * (int(seg.slot_id) + 1))) self.num_slots = len(self.member_names) def _fold_slot(self, batch: Dict[str, torch.Tensor], index: int) -> torch.Tensor: @@ -110,7 +116,9 @@ def _fold_slot(self, batch: Dict[str, torch.Tensor], index: int) -> torch.Tensor else: lengths = batch[members[0] + ".lengths"].to(torch.int64) num_holes = int(torch.sum(lengths)) if is_sequence else int(lengths.numel()) - keys = torch.zeros(num_holes, dtype=torch.int64, device=first.device) + # seeded with the salt, not zero: a hole no value reaches -- an empty + # DEEP member, a zero-length item -- is still that slot's hole + keys = torch.full((num_holes,), salt, dtype=torch.int64, device=first.device) for member_index in range(len(members)): member = members[member_index] raw = batch[member + ".values"] diff --git a/tzrec/prompt/hole_keys_test.py b/tzrec/prompt/hole_keys_test.py index 717cb7c5..15d4afdc 100644 --- a/tzrec/prompt/hole_keys_test.py +++ b/tzrec/prompt/hole_keys_test.py @@ -145,6 +145,24 @@ def test_two_slots_holding_the_same_id_do_not_collide(self) -> None: other = HoleKeyBuilder(_plan((_slot("beh", slot_id=1),))) self.assertFalse(torch.equal(self.module(batch), other(batch))) + def test_two_empty_holes_do_not_collide(self) -> None: + """A hole no value reaches still carries its slot: the salt seeds it.""" + batch = _tensors( + { + "beh.values": np.array([], dtype=np.int64), + "beh.lengths": np.array([0], dtype=np.int64), + } + ) + deep = HoleKeyBuilder( + _plan((_slot("beh", group_type=FeatureGroupType.DEEP, slot_id=0),)) + ) + other = HoleKeyBuilder( + _plan((_slot("beh", group_type=FeatureGroupType.DEEP, slot_id=1),)) + ) + self.assertEqual(deep(batch).numel(), 1) + self.assertNotEqual(int(deep(batch)[0]), 0) + self.assertNotEqual(int(deep(batch)[0]), int(other(batch)[0])) + def test_a_permuted_multi_value_item_does_not_collide(self) -> None: """``[a, b, c]`` and ``[c, b, a]`` are different items, in different bands.""" From 7679d2c86bcffab688e3caa3470e22d1970d08e3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:05:10 +0800 Subject: [PATCH 18/21] [refactor] expose compiled_prompt on the genrec front-end only ScriptWrapper assembles the prompt for any module carrying compiled_prompt, so the base model advertising it made every wrapper build a serving assembler and fold it never uses. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/models/genrec_model.py | 7 +------ tzrec/models/model.py | 8 ++++---- 2 files changed, 5 insertions(+), 10 deletions(-) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index ad91e7d2..ebe15408 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -191,11 +191,6 @@ def hf_backbone(self) -> nn.Module: """The HF module export and checkpointing reach for.""" return self.lm - @property - def compiled_prompt(self) -> CompiledPrompt: - """The prompt this model was built against.""" - return self._prompt - def build_input(self, batch: Batch) -> torch.Tensor: """Build packed LM input embeddings and fill projected positions. @@ -362,7 +357,7 @@ def __init__(self, model: BaseGenRecModel) -> None: self.embedding_group = model.embedding_group self.projections = model.projections self._slot_projections = list(model._slot_projections) - self._prompt = model.compiled_prompt + self._prompt = model._prompt self._features = list(model.features) self._hidden_size = int(model.lm.config.hidden_size) diff --git a/tzrec/models/model.py b/tzrec/models/model.py index 02636d2e..e14e7acf 100644 --- a/tzrec/models/model.py +++ b/tzrec/models/model.py @@ -391,10 +391,10 @@ def forward( class ScriptWrapper(BaseModule): """Model inference wrapper for jit.script. - A module that exposes ``compiled_prompt`` gets its prompt assembled here, - from the parsed dict, exactly as the training collator does: one walk, two - call sites. Serving alone also folds the prefix-cache ``hole_keys`` here, - which training never computes. + A module that exposes ``compiled_prompt`` -- the genrec front-end -- gets + its prompt assembled here from the parsed dict, exactly as the training + collator does: one walk, two call sites. Serving alone also folds the + prefix-cache ``hole_keys`` here, which training never computes. """ def __init__(self, module: nn.Module) -> None: From 42d8eeebfbbb374fc0abc8257a35af6da8e59c0c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:05:12 +0800 Subject: [PATCH 19/21] [doc] correct the prompt comments the check removal left stale max_length is a compile-time ceiling now, the sparse export filters projection heads rather than a backbone, and the tokenizer config names tokens, not ids. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/prompt/compile.py | 2 +- tzrec/prompt/types.py | 3 ++- tzrec/protos/prompt.proto | 5 +++-- tzrec/utils/export_util.py | 3 ++- 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/tzrec/prompt/compile.py b/tzrec/prompt/compile.py index f28666bc..492bbf0f 100644 --- a/tzrec/prompt/compile.py +++ b/tzrec/prompt/compile.py @@ -222,7 +222,7 @@ def _save_tokenizer_dir( ``tokenizer.json`` carries the vocabulary and the added SID atoms; the minimal ``tokenizer_config.json`` beside it names the tokenizer class and - the two special ids the prompt resolved, which is all a serving runtime + the two special tokens the prompt resolved, which is all a serving runtime needs to decode a generated SID atom through ``--tokenizer-path``. """ os.makedirs(tokenizer_dir, exist_ok=True) diff --git a/tzrec/prompt/types.py b/tzrec/prompt/types.py index b8d5aa0d..d50b4405 100644 --- a/tzrec/prompt/types.py +++ b/tzrec/prompt/types.py @@ -142,7 +142,8 @@ class PromptPlan: Args: segments: prompt body, in emission order. response_segments: supervised tail, in emission order. - max_length: validation ceiling; an over-long row is an error. + max_length: compile-time ceiling; compile refuses a plan whose + proven maximum exceeds it. Rows are not measured at runtime. max_total_length: proven ceiling when every slot is bounded, else None. max_holes: per-row projected-position ceiling, not a runtime shape. logits_suffix_len: upper bound on the supervised logits window. diff --git a/tzrec/protos/prompt.proto b/tzrec/protos/prompt.proto index 29d41f08..d8f23e6f 100644 --- a/tzrec/protos/prompt.proto +++ b/tzrec/protos/prompt.proto @@ -26,8 +26,9 @@ message PromptConfig { // does not enforce required, so compile_prompt raises the actual error. required SidSpace sid_space = 5; - // Validation ceiling, not a truncation trigger: an over-long row is an - // error, never truncated. + // Compile-time ceiling, not a truncation trigger: compile_prompt refuses a + // plan whose proven maximum length exceeds it. Rows are neither measured + // nor truncated at runtime. optional uint32 max_length = 6 [default = 0]; // Reserves a position filled by a projected slot. Materialized only when diff --git a/tzrec/utils/export_util.py b/tzrec/utils/export_util.py index fbc9400d..02d3ecdf 100644 --- a/tzrec/utils/export_util.py +++ b/tzrec/utils/export_util.py @@ -2688,7 +2688,8 @@ def _add_sparse_table( continue table_fqn = name[: -len(".weight")] export_emb_name = checkpoint_util.remap_input_tile_user_key(table_fqn) - # a dense weight beside the tables, such as a genrec backbone's + # a dense weight beside the tables, such as a genrec slot projection + # head, which the sparse export does not carry if export_emb_name not in emb_name_to_emb_dim: continue state_values_by_emb[export_emb_name] = values From 3543a81200a4e0ca0f66075c1a23452e0e97ca7a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:05:13 +0800 Subject: [PATCH 20/21] [ci] cover the multi-slot projection and compare token streams exactly The scripted round trip compared int64 streams through float, no front-end case had more than one projected slot, and the integration test's projected parameter was never false. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/models/genrec_model_test.py | 36 ++++++++++++++++++++++++-- tzrec/tests/genrec_integration_test.py | 17 ++++++------ tzrec/tests/utils.py | 6 ++--- 3 files changed, 44 insertions(+), 15 deletions(-) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 62088bd2..fb6dd618 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -27,7 +27,12 @@ project_slots, ) from tzrec.models.model import ScriptWrapper, TrainWrapper -from tzrec.prompt.assembler import HOLE_POSITIONS, INPUT_IDS, PromptAssembler +from tzrec.prompt.assembler import ( + HOLE_POSITIONS, + HOLE_SLOT_COUNTS, + INPUT_IDS, + PromptAssembler, +) from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder from tzrec.protos import feature_pb2 from tzrec.protos.model_pb2 import ModelConfig @@ -284,6 +289,30 @@ def test_front_end_returns_the_walk_and_the_projected_slots(self) -> None: self.assertTrue(torch.allclose(out[SLOT_EMBEDS], expected)) self.assertEqual(tuple(out[SLOT_EMBEDS].shape), (2, 32)) + def test_front_end_joins_several_projected_slots_in_order(self) -> None: + """Two slots: ``slot_embeds`` is their projections, occurrence by occurrence.""" + model, compiled = create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("beh", 8), _projected("ctx", 8)], + prompt="History : {{hist}} . {{beh}} then {{ctx}} Predict :", + ) + init_parameters(model, device=torch.device("cpu")) + wrapped = ScriptWrapper(GenRecFrontEnd(model)) + data = dict(self.data) + data["ctx.values"] = torch.tensor([1, 4, 7]) + data["ctx.lengths"] = torch.tensor([3]) + + out = wrapped(data) + counts = out[HOLE_SLOT_COUNTS] + self.assertEqual(counts.tolist(), [2, 3]) + grouped = model.embedding_group(wrapped.get_batch(data)) + spans = torch.split(out[SLOT_EMBEDS], counts.tolist()) + for seg, proj, span in zip( + compiled.prompt_plan.projected_slots, model._slot_projections, spans + ): + expected = proj(grouped[seg.name + seg.output_key]).reshape(-1, 32) + self.assertTrue(torch.allclose(span, expected), seg.name) + def test_front_end_shares_the_checkpoint_names_and_not_the_lm(self) -> None: names = set(self.wrapped.state_dict()) self.assertTrue(any(n.startswith("model.embedding_group.") for n in names)) @@ -296,7 +325,10 @@ def test_front_end_traces_and_scripts(self) -> None: scripted = torch.jit.script(symbolic_trace(self.wrapped)) out = scripted(self.data) for key, value in eager.items(): - self.assertTrue(torch.allclose(out[key].float(), value.float()), key) + if value.is_floating_point(): + self.assertTrue(torch.allclose(out[key], value), key) + else: + self.assertTrue(torch.equal(out[key], value), key) if __name__ == "__main__": diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index a69b25c1..6f63204a 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -65,7 +65,7 @@ def tearDown(self): if self.success and os.path.exists(self.test_dir): shutil.rmtree(self.test_dir) - def _prepare_config(self, projected: bool) -> str: + def _prepare_config(self) -> str: """Write the tiny backbone, tokenizer, manifest and data; return the config.""" backbone = os.path.join(self.test_dir, "backbone") create_tiny_causal_lm(64).save_pretrained(backbone) @@ -76,7 +76,7 @@ def _prepare_config(self, projected: bool) -> str: with open(manifest, "w") as f: json.dump({"codebook": _CODEBOOK, "bundle_uuid": _BUNDLE_UUID}, f) self.data_glob = utils.create_mock_prompt_data( - os.path.join(self.test_dir, "data"), _CODEBOOK, projected=projected + os.path.join(self.test_dir, "data"), _CODEBOOK ) config = config_util.load_pipeline_config(_MOCK_CONFIG) @@ -85,16 +85,15 @@ def _prepare_config(self, projected: bool) -> str: config.model_config.genrec_causal_lm_model.hf_model_name_or_path = backbone config.prompt_config.tokenizer_path = tokenizer config.prompt_config.sid_space.manifest_path = manifest - if projected: - text_format.Merge(_BEH, config.feature_configs.add()) - config.prompt_config.prompt = "History : {{hist}} . {{beh}} Predict :" + text_format.Merge(_BEH, config.feature_configs.add()) + config.prompt_config.prompt = "History : {{hist}} . {{beh}} Predict :" config_path = os.path.join(self.test_dir, "genrec.config") config_util.save_message(config, config_path) return config_path - def _train_eval_export(self, projected: bool) -> str: + def _train_eval_export(self) -> str: """Run the pipeline; return the trained ``pipeline.config`` path.""" - config_path = self._prepare_config(projected) + config_path = self._prepare_config() self.success = utils.test_train_eval(config_path, self.test_dir) trained = os.path.join(self.test_dir, "pipeline.config") if self.success: @@ -122,7 +121,7 @@ def _request(self, columns, rows: int = 4): return out def test_genrec_train_eval_export(self): - trained = self._train_eval_export(projected=True) + trained = self._train_eval_export() export_dir = os.path.join(self.test_dir, "export") for name in ( "scripted_model.pt", @@ -198,7 +197,7 @@ def test_genrec_train_eval_export(self): @unittest.skipIf(*gpu_unavailable) @mark_ci_scope("gpu") def test_genrec_export_distributed_embedding(self): - trained = self._train_eval_export(projected=True) + trained = self._train_eval_export() dist_dir = os.path.join(self.test_dir, "export_dist") self.success = utils.test_export( trained, diff --git a/tzrec/tests/utils.py b/tzrec/tests/utils.py index 06c97a1e..412a9427 100644 --- a/tzrec/tests/utils.py +++ b/tzrec/tests/utils.py @@ -535,7 +535,7 @@ def create_mock_data( def create_mock_prompt_data( - path: str, codebook: Sequence[int], projected: bool, num_rows: int = 8 + path: str, codebook: Sequence[int], num_rows: int = 8 ) -> str: """Write a parquet of SID histories for the prompt-native tests. @@ -545,7 +545,6 @@ def create_mock_prompt_data( Args: path: directory to write into. codebook: per-level SID vocabulary sizes. - projected: whether to add the ``beh`` column. num_rows: samples to write. Returns: @@ -561,9 +560,8 @@ def codes(items: int) -> List[int]: columns = { "hist": [codes(2) for _ in range(num_rows)], "answer": [codes(1) for _ in range(num_rows)], + "beh": [rng.integers(0, 32, size=2).tolist() for _ in range(num_rows)], } - if projected: - columns["beh"] = [rng.integers(0, 32, size=2).tolist() for _ in range(num_rows)] os.makedirs(path, exist_ok=True) pq.write_table( pa.table( From e5acb93b54b78e90a5a8cfe4958c7bb908588f3c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Tue, 8 Sep 2026 19:14:27 +0800 Subject: [PATCH 21/21] [doc] tighten the ScriptWrapper and HoleKeyBuilder docstrings ScriptWrapper wraps every model, so it leads with parsing a request into a batch and treats the prompt as the conditional it is; the fold's rationale reads in one paragraph. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01MB8aCo5mgF5ohAcBpWYDaN --- tzrec/models/model.py | 8 ++++---- tzrec/prompt/hole_keys.py | 18 +++++++----------- 2 files changed, 11 insertions(+), 15 deletions(-) diff --git a/tzrec/models/model.py b/tzrec/models/model.py index e14e7acf..1fe16e2f 100644 --- a/tzrec/models/model.py +++ b/tzrec/models/model.py @@ -391,10 +391,10 @@ def forward( class ScriptWrapper(BaseModule): """Model inference wrapper for jit.script. - A module that exposes ``compiled_prompt`` -- the genrec front-end -- gets - its prompt assembled here from the parsed dict, exactly as the training - collator does: one walk, two call sites. Serving alone also folds the - prefix-cache ``hole_keys`` here, which training never computes. + Parses a request dict into a ``Batch`` and runs the wrapped module's + ``predict``. A module that also exposes ``compiled_prompt`` has its prompt + assembled here by the same walk the training collator runs, with the + serving-only ``hole_keys`` fold beside it, which training never computes. """ def __init__(self, module: nn.Module) -> None: diff --git a/tzrec/prompt/hole_keys.py b/tzrec/prompt/hole_keys.py index 00e06163..bec1545d 100644 --- a/tzrec/prompt/hole_keys.py +++ b/tzrec/prompt/hole_keys.py @@ -60,17 +60,13 @@ def _wrap64(value: int) -> int: class HoleKeyBuilder(nn.Module): """Folds every projected slot's input values into one key per hole. - Every member value that produces a hole contributes, discriminated by - slot, by member and by its index inside the hole. Without the first, two - slots holding the same id would match; without the last two, a two-member - slot with values ``(a, b)`` would match one with ``(b, a)`` and a permuted - multi-value item would match itself reordered -- all plausible, all wrong, - and all silent. A dense member contributes its float32 bit pattern per row, - which is the parsed input and not a computed reduction. - - Body and response slots are walked together, which the assembler's holes - match because compile gives every response slot ``FillMode.INLINE``: a - response slot is never PROJECTED and so never reaches ``projected_slots``. + Every value contributing to a hole is discriminated by slot, by member and + by its index inside the hole; without those, two slots holding the same id, + a two-member slot with ``(a, b)`` against ``(b, a)``, and a permuted + multi-value item would all match silently. A dense member contributes its + float32 bit pattern per row, the parsed input rather than a computed + reduction. Response slots compile INLINE, so ``projected_slots`` is + body-only and stays aligned with the assembler's holes. Args: prompt_plan: the compiled plan; its ``projected_slots`` fix the hole