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)
|
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)
|
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 = {}
|