Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
79e7f69
[feat] export the prompt front-end and serving contract for genrec
tiankongdeguiji Sep 7, 2026
b307996
[refactor] make the prompt assembler one scripted module with two cal…
tiankongdeguiji Sep 7, 2026
40169c3
[refactor] export the genrec front-end through export_model
tiankongdeguiji Sep 7, 2026
17504f6
[refactor] keep only the SID space in prompt.json
tiankongdeguiji Sep 7, 2026
6f5869f
[refactor] move the hole_keys fold out of the prompt assembler
tiankongdeguiji Sep 7, 2026
3a261c7
[refactor] drop the prompt digests and the restore guard
tiankongdeguiji Sep 7, 2026
ef71e5d
[refactor] rename Genrec to GenRec
tiankongdeguiji Sep 8, 2026
db61bbc
[refactor] write the HF assets from export, not from the front-end
tiankongdeguiji Sep 8, 2026
6f7ea37
[ci] run the genrec pipeline as an integration test
tiankongdeguiji Sep 8, 2026
5e45168
[ci] mark the GPU-only hole-key fold test with its CI scope
tiankongdeguiji Sep 8, 2026
c5fc3af
[doc] shorten the genrec module docstrings
tiankongdeguiji Sep 8, 2026
03a4e56
[refactor] trust the parsed batch in the prompt walk
tiankongdeguiji Sep 8, 2026
044d2bd
[refactor] fold the genrec test fixtures into test_util
tiankongdeguiji Sep 8, 2026
c88e2da
[refactor] write one config.json for the genrec export
tiankongdeguiji Sep 8, 2026
5a169f0
[chore] bump version to 1.4.5
tiankongdeguiji Sep 8, 2026
a59f669
Merge remote-tracking branch 'origin/master' into feat/genrec-serving…
tiankongdeguiji Sep 8, 2026
48adf07
[perf] load only the backbone at genrec export
tiankongdeguiji Sep 8, 2026
a49fda2
[refactor] seed each hole key with its slot salt
tiankongdeguiji Sep 8, 2026
7679d2c
[refactor] expose compiled_prompt on the genrec front-end only
tiankongdeguiji Sep 8, 2026
42d8eee
[doc] correct the prompt comments the check removal left stale
tiankongdeguiji Sep 8, 2026
3543a81
[ci] cover the multi-slot projection and compare token streams exactly
tiankongdeguiji Sep 8, 2026
e5acb93
[doc] tighten the ScriptWrapper and HoleKeyBuilder docstrings
tiankongdeguiji Sep 8, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 3 additions & 8 deletions tzrec/datasets/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -41,7 +40,7 @@
remove_nullable,
)
from tzrec.features.feature import BaseFeature
from tzrec.prompt.assembler import 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
Expand Down Expand Up @@ -401,12 +400,8 @@ 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:
batch.additional_infos.update(
{
k: torch.from_numpy(np.asarray(v))
for k, v in self._prompt_assembler.forward(output_data).items()
}
)
streams = self._prompt_assembler(output_data)
batch.additional_infos.update({k: streams[k] for k in OUTPUT_KEYS})

# Set checkpoint info on batch
batch.checkpoint_info = checkpoint_info
Expand Down
85 changes: 32 additions & 53 deletions tzrec/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -102,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
Expand Down Expand Up @@ -836,8 +836,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:
Expand Down Expand Up @@ -1122,7 +1120,6 @@ def evaluate(
)

if checkpoint_path:
check_prompt_assets(compiled_prompt, checkpoint_path)
ckpt_manager.restore(
checkpoint_path,
model,
Expand Down Expand Up @@ -1207,64 +1204,20 @@ 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 <checkpoint>/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),
)
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,
)
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)

# 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
Expand Down Expand Up @@ -1312,6 +1265,33 @@ 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
if config_util.use_dense_ema(
pipeline_config.export_config, pipeline_config.train_config
):
raise ValueError(
"HF export: dcp_to_hf reads <checkpoint>/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.")
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,
front_end,
checkpoint_path,
export_dir,
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,
Expand Down Expand Up @@ -1789,7 +1769,6 @@ def predict_checkpoint(
model.eval()

if checkpoint_path:
check_prompt_assets(compiled_prompt, checkpoint_path)
ckpt_manager.restore(
checkpoint_path,
model,
Expand Down
19 changes: 0 additions & 19 deletions tzrec/main_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -366,23 +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 <ckpt>/model unconditionally, so it would silently
# ship raw weights where TorchScript export ships the EMA ones.
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"))


class PredictionLifecycleTest(unittest.TestCase):
"""Tests for prediction lifecycle wiring."""
Expand Down
28 changes: 14 additions & 14 deletions tzrec/models/genrec_causal_lm_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,20 +26,20 @@

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,
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
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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand All @@ -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,
Expand All @@ -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.

Expand All @@ -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.

Expand Down
Loading
Loading