| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
|
|
| from easydict import EasyDict as edict |
|
|
|
|
| @dataclass |
| class SampleOutputBase: |
| true_local_ids_map: dict[str, dict[str, list[int]]] |
| accepted_ids_map: dict[str, dict[str, list[int]]] |
| sampled_tokens_map: dict[str, dict[str, list[int]]] |
| mask_token_rel_ids_map: dict[str, dict[str, list[int]]] | None = None |
| confidence_map: dict[str, dict[str, list[float]]] | None = None |
| initial_confidence_map: dict[str, dict[str, list[float]]] | None = None |
| edit_writes_map: dict[str, dict[str, dict[int, int]]] | None = None |
| block_state_map: dict[str, dict[str, dict]] | None = None |
|
|
| def __post_init__(self): |
| req_ids = set(self.accepted_ids_map.keys()) |
| self.accepted_ids_map = edict(self.accepted_ids_map) |
| self.sampled_tokens_map = edict(self.sampled_tokens_map) |
| self.true_local_ids_map = edict(self.true_local_ids_map) |
| self.mask_token_rel_ids_map = edict(self.mask_token_rel_ids_map or {}) |
| self.confidence_map = edict(self.confidence_map or {}) |
| self.initial_confidence_map = edict(self.initial_confidence_map or {}) |
| edit_writes_map = self.edit_writes_map or {} |
| block_state_map = self.block_state_map or {} |
| for req_id_str in req_ids: |
| edit_writes_map.setdefault(req_id_str, {}) |
| block_state_map.setdefault(req_id_str, {}) |
| self.edit_writes_map = edict(edit_writes_map) |
| self.block_state_map = edict(block_state_map) |
|
|
|
|
| def merge_sample_outputs(outputs: list[SampleOutputBase | None]) -> SampleOutputBase: |
| true_local_ids_map: dict[str, dict[str, list[int]]] = {} |
| accepted_ids_map: dict[str, dict[str, list[int]]] = {} |
| sampled_tokens_map: dict[str, dict[str, list[int]]] = {} |
| mask_token_rel_ids_map: dict[str, dict[str, list[int]]] = {} |
| confidence_map: dict[str, dict[str, list[float]]] = {} |
| initial_confidence_map: dict[str, dict[str, list[float]]] = {} |
| edit_writes_map: dict[str, dict[str, dict[int, int]]] = {} |
| block_state_map: dict[str, dict[str, dict]] = {} |
|
|
| for output in outputs: |
| if output is None: |
| continue |
| true_local_ids_map.update(dict(output.true_local_ids_map)) |
| accepted_ids_map.update(dict(output.accepted_ids_map)) |
| sampled_tokens_map.update(dict(output.sampled_tokens_map)) |
| mask_token_rel_ids_map.update(dict(output.mask_token_rel_ids_map)) |
| confidence_map.update(dict(output.confidence_map)) |
| initial_confidence_map.update(dict(output.initial_confidence_map)) |
| edit_writes_map.update(dict(output.edit_writes_map)) |
| block_state_map.update(dict(output.block_state_map)) |
|
|
| return SampleOutputBase( |
| true_local_ids_map=true_local_ids_map, |
| accepted_ids_map=accepted_ids_map, |
| sampled_tokens_map=sampled_tokens_map, |
| mask_token_rel_ids_map=mask_token_rel_ids_map, |
| confidence_map=confidence_map, |
| initial_confidence_map=initial_confidence_map, |
| edit_writes_map=edit_writes_map, |
| block_state_map=block_state_map, |
| ) |
|
|