Skip to content

guidellm.data.deserializers.trace_mooncake

The Mooncake trace format and data arguments.

Reads a trace file (timestamp, input_length, output_length, hash_ids) and yields one row per line with a synthetic prompt matching the requested input_length for replay benchmarks. Checks for distinctness between hash IDs that share the same previous hash ID.

MooncakeTraceFormat

Bases: SingleTurnTraceFormat

Mooncake trace format requires a column for timestamps, prompt token counts, ouput token counts and lists of hash IDs.

Hash IDs are globally unique identifiers based on the current and previous token blocks in a prompt. The relationships of IDs forms a tree, where every first ID in a prompt has a parent node of None. Parent nodes can have an unbounded number of children. Two hash IDs can represent identical blocks of tokens so long as they do not share the same parent (previous ID).

For more details, see section 4 of https://arxiv.org/pdf/2407.00079.

Generated prompts match the prompt token count of the row.

Source code in src/guidellm/data/deserializers/trace_mooncake.py
@TraceFormatRegistry.register("mooncake")
class MooncakeTraceFormat(SingleTurnTraceFormat):
    """Mooncake trace format requires a column for timestamps, prompt token counts,
    ouput token counts and lists of hash IDs.

    Hash IDs are globally unique identifiers based on the current and previous token
    blocks in a prompt. The relationships of IDs forms a tree, where every first ID
    in a prompt has a parent node of `None`. Parent nodes can have an unbounded
    number of children. Two hash IDs can represent identical blocks of tokens so long
    as they do not share the same parent (previous ID).

    For more details, see section 4 of https://arxiv.org/pdf/2407.00079.

    Generated prompts match the prompt token count of the row."""

    def __init__(self, config: MooncakeTraceFormatArgs, dataset: Dataset) -> None:
        self.config = config
        self.dataset = dataset
        self._hash_id_table: dict[int, tuple[int, ...]] = {}
        self._sibling_table: dict[Any, set[tuple[int, ...]]] = {}

    def reset_hash_tables(self) -> None:
        """Replace hash tables so this copy does not reuse earlier tokens."""
        self._hash_id_table = {}
        self._sibling_table = {}

    def required_columns(self) -> Features:
        return Features({self.config.hash_ids_column: List(Value("int32"))})

    def find_required_columns(self, columns: list[str]) -> list[str]:
        return get_missing_columns(columns, self.dataset.column_names)

    def validate_row(self, row: dict) -> None:
        n_in = row[self.config.prompt_tokens_column]
        n_blocks = len(row[self.config.hash_ids_column])
        block_size = self.config.hash_id_block_size
        for hash_id in row[self.config.hash_ids_column]:
            if hash_id < 0:
                raise InvalidRowError(f"Hash ID must be non-negative, got {hash_id}")
        if math.ceil(n_in / block_size) != n_blocks:
            raise InvalidRowError(
                f"Input token count of {n_in} split into blocks of size "
                f"{block_size} does not match given {n_blocks} blocks"
            )

    def create_prompt(
        self, row: dict, processor: PreTrainedTokenizerBase, faker: Faker
    ) -> str:
        """Before generating the prompt, this first generates a block of tokens for
        each hash ID that has not already been seen."""
        ids = row[self.config.hash_ids_column]
        fill_hash_id_table(
            ids,
            self._hash_id_table,
            self._sibling_table,
            processor,
            faker,
            lambda _idx, hash_id: _calculate_required_prompt_tokens(
                self.config, row, hash_id
            ),
        )
        return create_prompt_from_hash_ids(ids, self._hash_id_table, processor)

create_prompt(row, processor, faker)

Before generating the prompt, this first generates a block of tokens for each hash ID that has not already been seen.

Source code in src/guidellm/data/deserializers/trace_mooncake.py
def create_prompt(
    self, row: dict, processor: PreTrainedTokenizerBase, faker: Faker
) -> str:
    """Before generating the prompt, this first generates a block of tokens for
    each hash ID that has not already been seen."""
    ids = row[self.config.hash_ids_column]
    fill_hash_id_table(
        ids,
        self._hash_id_table,
        self._sibling_table,
        processor,
        faker,
        lambda _idx, hash_id: _calculate_required_prompt_tokens(
            self.config, row, hash_id
        ),
    )
    return create_prompt_from_hash_ids(ids, self._hash_id_table, processor)

reset_hash_tables()

Replace hash tables so this copy does not reuse earlier tokens.

Source code in src/guidellm/data/deserializers/trace_mooncake.py
def reset_hash_tables(self) -> None:
    """Replace hash tables so this copy does not reuse earlier tokens."""
    self._hash_id_table = {}
    self._sibling_table = {}