Download source/src/speculators/train/data.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 22.5 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/train/data.py
- Command line
-
hf download hf://khazic/spec-b300/source/src/speculators/train/data.py
-
curl -L -o data.py https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/train/data.py
22.5 kB
| import logging | |
| import warnings | |
| from collections.abc import Callable, Sequence | |
| from concurrent.futures import ThreadPoolExecutor | |
| from os import PathLike | |
| from pathlib import Path | |
| from typing import Any, Literal, cast | |
| import httpx | |
| import openai | |
| import torch | |
| from datasets import load_from_disk | |
| from torch.utils.data import Dataset | |
| from hs_connectors import FileTransfer, HiddenStatesTransfer | |
| from speculators.data_generation.offline import check_hidden_states | |
| from speculators.data_generation.vllm_client import ( | |
| DEFAULT_MAX_RETRIES, | |
| DEFAULT_REQUEST_TIMEOUT, | |
| ClientItem, | |
| generate_hidden_states, | |
| ) | |
| from speculators.train.noise_transforms import TransformTensors | |
| from speculators.train.recovery import ( | |
| RECOVERY_METADATA_KEY, | |
| GenerationRecoveryGuard, | |
| RecoveryMetadata, | |
| SampleUnavailable, | |
| ) | |
| BatchType = dict[str, Any] | |
| logger = logging.getLogger("speculators") | |
| def create_empty_sample( | |
| hidden_size: int, | |
| num_target_layers: int = 3, | |
| dtype: torch.dtype = torch.bfloat16, | |
| ): | |
| # data structure: { | |
| # "hidden_states": [seq_len, num_target_layers * hidden_size], | |
| # "input_ids": [seq_len], | |
| # "verifier_last_hidden_states": [seq_len, hidden_size], | |
| # "loss_mask": [seq_len], | |
| # "lengths": [1], | |
| # "position_ids": [seq_len], | |
| # } | |
| # Default dtype is bfloat16 to match the hidden_states dtype used downstream. | |
| # When this fallback is used (e.g. vLLM hidden-state extraction times out and | |
| # we substitute an empty sample), the implicit float32 placeholders crashed | |
| # bf16 EAGLE-3 layers (fc, verifier_lm_head) with a dtype mismatch. | |
| return { | |
| "hidden_states": torch.empty(0, num_target_layers * hidden_size, dtype=dtype), | |
| "input_ids": torch.empty(0, dtype=torch.long), | |
| "verifier_last_hidden_states": torch.empty(0, hidden_size, dtype=dtype), | |
| "loss_mask": torch.empty(0, dtype=torch.bool), | |
| "lengths": torch.tensor([0], dtype=torch.long), | |
| "position_ids": torch.arange(0, dtype=torch.long), | |
| } | |
| def _has_multimodal_content(messages: list[dict]) -> bool: | |
| """True when any turn carries non-text content (images, video, audio). | |
| Text-only turns store ``content`` as a plain string. Multimodal turns | |
| (produced by ``_adapt_conv_for_vllm``) store it as a list of typed parts, | |
| e.g. ``[{"type": "text", ...}, {"type": "image_url", ...}]``. | |
| """ | |
| return any(isinstance(m.get("content"), list) for m in messages) | |
| def build_client_item(dataset_item: dict) -> ClientItem: | |
| """Build a request payload for vLLM hidden-state extraction. | |
| When ``messages`` is included, ``generate_hidden_states`` uses the Chat | |
| Completions API and vLLM **re-tokenizes from the raw messages**, ignoring | |
| ``input_ids``. This is required for multimodal inputs (the Completions | |
| API cannot carry image/video/audio references), but harmful for text-only | |
| data: preprocessing truncates ``input_ids`` to ``seq_length``, yet the | |
| ``messages`` column stores the original un-truncated conversation. | |
| Re-tokenizing those messages produces a longer sequence that can exceed | |
| ``max_model_len``. | |
| We therefore only forward ``messages`` when the conversation actually | |
| contains multimodal content. Text-only conversations always go through | |
| the Completions API with the pre-truncated ``input_ids``. | |
| This matters for models like Qwen3.5-0.8B whose ``AutoProcessor`` returns | |
| a ``ProcessorMixin`` (``Qwen3VLProcessor``), causing preprocessing to | |
| populate the ``messages`` column even for purely text-only datasets. | |
| Text-only EAGLE-3 models (e.g. Llama) use a plain tokenizer, so | |
| ``messages`` is never created and this guard is a no-op. | |
| """ | |
| out_dict: dict = {"input_ids": dataset_item["input_ids"].tolist()} | |
| if "messages" in dataset_item and _has_multimodal_content(dataset_item["messages"]): | |
| out_dict["messages"] = dataset_item["messages"] | |
| return cast("ClientItem", out_dict) | |
| class BaseDataset(Dataset): | |
| def __init__( | |
| self, | |
| max_len: int, | |
| transform: TransformTensors | None = None, | |
| hidden_states_dtype=torch.bfloat16, | |
| fetch_threads: int = 1, | |
| ): | |
| self.max_len = max_len | |
| self.transform = transform | |
| self.hidden_states_dtype = hidden_states_dtype | |
| self.fetch_threads = max(1, fetch_threads) | |
| self._fetch_pool: ThreadPoolExecutor | None = None | |
| self.approx_lengths = self._compute_approx_lengths() | |
| def _compute_approx_lengths(self): | |
| raise NotImplementedError | |
| def _get_raw_data(self, index: int) -> BatchType | SampleUnavailable: | |
| raise NotImplementedError | |
| def _prepare_fetch(self) -> None: | |
| """One-time setup that must happen before the fetch pool starts. | |
| Subclasses with lazily initialized state override this so several fetch | |
| threads don't race to build it. | |
| """ | |
| def __getstate__(self) -> dict[str, Any]: | |
| # The DataLoader pickles this dataset into every spawned worker, and a | |
| # live executor cannot cross that boundary. Workers rebuild their own | |
| # pool on first use. | |
| state = self.__dict__.copy() | |
| state["_fetch_pool"] = None | |
| return state | |
| def __getitems__( | |
| self, indices: Sequence[int] | |
| ) -> list[BatchType | SampleUnavailable]: | |
| """Fetch a whole batch at once, overlapping the per-sample round trips. | |
| ``_MapDatasetFetcher`` prefers this over ``__getitem__`` when it exists, | |
| handing over every index in the batch together. Online hidden states | |
| cost one blocking HTTP request per sample, so fetching serially makes a | |
| step wait on the *sum* of its samples' latencies -- and because the | |
| DataLoader delivers batches in order, one slow sample stalls the whole | |
| rank while its peers idle at the gradient all-reduce. Threads turn that | |
| sum into a max. | |
| The returned order matches ``indices``: collation derives ``lengths`` | |
| and ``document_ids`` from the sequence position. | |
| """ | |
| if self.fetch_threads <= 1 or len(indices) <= 1: | |
| return [self[index] for index in indices] | |
| self._prepare_fetch() | |
| if self._fetch_pool is None: | |
| self._fetch_pool = ThreadPoolExecutor( | |
| max_workers=self.fetch_threads, | |
| thread_name_prefix="hs-fetch", | |
| ) | |
| # map() preserves input order and re-raises unexpected exceptions from the | |
| # threads; ordinary generation failures are reported through recovery metadata. | |
| return list(self._fetch_pool.map(self.__getitem__, indices)) | |
| def __getitem__(self, index) -> BatchType | SampleUnavailable: | |
| data = self._get_raw_data(index) | |
| if isinstance(data, SampleUnavailable): | |
| return data | |
| # data structure: { | |
| # "hidden_states": [seq_len, 3 * hidden_size], | |
| # "input_ids": [seq_len], | |
| # "verifier_last_hidden_states": [seq_len, hidden_size], | |
| # "loss_mask": [seq_len], | |
| # } | |
| # Add lengths tensor | |
| seq_len = data["input_ids"].shape[0] | |
| data["lengths"] = torch.tensor([seq_len], dtype=torch.long) | |
| # shape: [1] | |
| data.setdefault("position_ids", torch.arange(seq_len, dtype=torch.long)) | |
| # shape: [seq_len] | |
| # data structure: { | |
| # "hidden_states": [seq_len, 3 * hidden_size], | |
| # "input_ids": [seq_len], | |
| # "verifier_last_hidden_states": [seq_len, hidden_size], | |
| # "loss_mask": [seq_len], | |
| # "lengths": [1], | |
| # "position_ids": [seq_len], | |
| # } | |
| # Apply transform | |
| if self.transform: | |
| data = self.transform(data) | |
| return data | |
| class ArrowDataset(BaseDataset): | |
| def __init__( | |
| self, | |
| max_len: int, | |
| datapath: str | PathLike, | |
| transfer: HiddenStatesTransfer | None = None, | |
| vllm_endpoint: str = "http://localhost:8000/v1", | |
| on_missing: Literal["generate", "skip", "warn", "raise"] = "generate", | |
| on_generate: Literal["cache", "delete"] = "delete", | |
| train_ratio: float = 1.0, | |
| split: Literal["train", "val"] = "train", | |
| transform: TransformTensors | None = None, | |
| hidden_states_dtype=torch.bfloat16, | |
| model: str | None = None, | |
| request_timeout: float | None = DEFAULT_REQUEST_TIMEOUT, | |
| max_retries: int = DEFAULT_MAX_RETRIES, | |
| generation_validation_retries: int = 2, | |
| max_consecutive_generation_failures: int = 20, | |
| fail_on_hidden_state_error: bool = False, | |
| fetch_threads: int = 1, | |
| http_keepalive: bool = True, | |
| ): | |
| self.data = load_from_disk(datapath) | |
| if not 0.0 < train_ratio <= 1.0: | |
| raise ValueError(f"train_ratio must be in (0.0, 1.0], got {train_ratio}") | |
| if split == "val" and train_ratio == 1.0: | |
| raise ValueError("train_ratio=1.0 leaves no validation split") | |
| # Both splits derive their boundary from this one expression, | |
| # so they are exactly complementary. | |
| split_idx = int(len(self.data) * train_ratio) | |
| start, stop = ( | |
| (0, split_idx) if split == "train" else (split_idx, len(self.data)) | |
| ) | |
| if start >= stop: | |
| raise ValueError( | |
| f"{split} split is empty (dataset has {len(self.data)} rows, " | |
| f"train_ratio={train_ratio} gives split_idx={split_idx})" | |
| ) | |
| self.start_file_idx = start | |
| self.data = self.data.select(range(start, stop)) | |
| self.transfer = transfer or FileTransfer(Path(datapath) / "hidden_states") | |
| self.vllm_endpoint = vllm_endpoint | |
| self.on_missing = on_missing | |
| self.on_generate = on_generate | |
| self.client: openai.OpenAI | None = None | |
| self.model = model | |
| self.request_timeout = request_timeout | |
| self.max_retries = max_retries | |
| self.fail_on_hidden_state_error = fail_on_hidden_state_error | |
| self.http_keepalive = http_keepalive | |
| self.generation_recovery = GenerationRecoveryGuard( | |
| retries=0 if fail_on_hidden_state_error else generation_validation_retries, | |
| max_consecutive_failures=( | |
| 1 | |
| if fail_on_hidden_state_error | |
| else max_consecutive_generation_failures | |
| ), | |
| ) | |
| # Delay super init so that `_compute_approx_lengths` has required data | |
| super().__init__(max_len, transform, hidden_states_dtype, fetch_threads) | |
| def __getstate__(self) -> dict[str, Any]: | |
| # An openai/httpx client owns sockets that cannot cross the spawn boundary | |
| # into a DataLoader worker; each worker builds its own in `_prepare_fetch`. | |
| state = super().__getstate__() | |
| state["client"] = None | |
| return state | |
| def _map_to_file_idx(self, index: int): | |
| return index + self.start_file_idx | |
| def _prepare_fetch(self) -> None: | |
| """Build the vLLM client before the fetch pool starts, not inside it. | |
| ``_get_raw_data`` sets the client up on first use. Several fetch threads | |
| discovering ``self.client is None`` at once would each run the model | |
| handshake and ``transfer.setup()``. | |
| """ | |
| if not self.client: | |
| self._setup_client() | |
| def _setup_client(self): | |
| http_client = None | |
| if not self.http_keepalive: | |
| # vLLM hands every connection to one of its --api-server-count frontend | |
| # processes and keeps it there, and each of those runs the multimodal | |
| # processor on a single thread (vllm/renderers/base.py builds | |
| # ``_mm_executor`` with max_workers=1). A pooled connection therefore | |
| # pins this worker's requests to one CPU slot no matter how wide the | |
| # frontend is; retiring the connection after each response sends the | |
| # next request through a fresh accept() and spreads the load. | |
| http_client = httpx.Client( | |
| limits=httpx.Limits(max_keepalive_connections=0) | |
| ) | |
| client = openai.OpenAI( | |
| base_url=self.vllm_endpoint, | |
| api_key="EMPTY", | |
| max_retries=0, | |
| http_client=http_client, | |
| ) | |
| list_models = client.models.list() | |
| model_id = list_models.data[0].id | |
| if self.model and self.model != model_id: | |
| raise ValueError( | |
| f"An explicit model name was passed ({self.model}) which doesn't match" | |
| f" found model_id {model_id}." | |
| "Please make sure --endpoint is set to the correct vllm instance." | |
| ) | |
| self.model = model_id | |
| self.transfer.setup() | |
| # Do not retain a half-initialized client if model discovery or Mooncake | |
| # setup failed; the outer full-round-trip retry should redo initialization. | |
| self.client = client | |
| def __len__(self): | |
| return len(self.data) | |
| def _compute_approx_lengths(self) -> list[int]: | |
| """Get lengths of the dataset samples.""" | |
| return list(self.data.with_format(None)["seq_len"]) | |
| def _generate_hidden_states_once( | |
| self, | |
| index: int, | |
| dataset_item: dict, | |
| client_item: ClientItem, | |
| ) -> dict[str, torch.Tensor]: | |
| handle: str | None = None | |
| try: | |
| if not self.client: | |
| self._setup_client() | |
| handle = generate_hidden_states( | |
| self.client, # type:ignore[arg-type] | |
| self.model, # type:ignore[arg-type] | |
| client_item, | |
| timeout=self.request_timeout, | |
| max_retries=self.max_retries, | |
| ) | |
| loaded_hs = self.transfer.get_generated(handle) | |
| if loaded_hs is None: | |
| raise ValueError(f"Failed to load hidden states for handle {handle}") | |
| # Covers token/shape mismatches and non-finite values. The Mooncake | |
| # transfer performs manifest/checksum validation first. | |
| check_hidden_states(loaded_hs, dataset_item["input_ids"].tolist()) | |
| file_idx = self._map_to_file_idx(index) | |
| if self.on_generate == "cache": | |
| self.transfer.cache(handle, file_idx) | |
| else: | |
| try: | |
| self.transfer.delete(handle) | |
| except Exception as cleanup_error: # noqa: BLE001 | |
| logger.warning( | |
| "Loaded a valid hidden-state sample but failed to delete " | |
| "handle %s: %s", | |
| handle, | |
| cleanup_error, | |
| ) | |
| return loaded_hs | |
| except Exception: | |
| if handle is not None: | |
| try: | |
| self.transfer.delete(handle) | |
| except Exception as cleanup_error: # noqa: BLE001 | |
| logger.warning( | |
| "Failed to clean generated hidden-state handle %s: %s", | |
| handle, | |
| cleanup_error, | |
| ) | |
| raise | |
| def _get_raw_data(self, index: int) -> BatchType | SampleUnavailable: | |
| file_idx = self._map_to_file_idx(index) | |
| cached_hs = self.transfer.get_cached(file_idx) | |
| if cached_hs is None: | |
| if self.on_missing == "generate": | |
| dataset_item = self.data[index] | |
| client_item = build_client_item(dataset_item) | |
| loaded_hs = self.generation_recovery.run( | |
| lambda: self._generate_hidden_states_once( | |
| index, | |
| dataset_item, | |
| client_item, | |
| ), | |
| description=( | |
| f"Hidden-state round trip failed for dataset index {index}, " | |
| f"file index {file_idx}" | |
| ), | |
| ) | |
| elif self.on_missing == "skip": | |
| return SampleUnavailable() | |
| elif self.on_missing == "warn": | |
| warnings.warn( | |
| f"Failed to load hidden states for sample {index}. Skipping...", | |
| stacklevel=1, | |
| ) | |
| return SampleUnavailable( | |
| reason=f"Hidden states unavailable for sample {index}" | |
| ) | |
| else: | |
| raise RuntimeError(f"Failed to load hidden states for sample {index}.") | |
| else: | |
| loaded_hs = cached_hs | |
| if isinstance(loaded_hs, SampleUnavailable): | |
| return loaded_hs | |
| # loaded_hs structure: { | |
| # "hidden_states": [seq_len, num_layers, hidden_size] | |
| # "token_ids": [seq_len] | |
| # } | |
| if not torch.equal(loaded_hs["token_ids"], self.data[index]["input_ids"]): | |
| if self.fail_on_hidden_state_error: | |
| raise RuntimeError( | |
| f"Loaded hidden-state token ids do not match sample {index}" | |
| ) | |
| warnings.warn( | |
| f"Loaded token ids {loaded_hs['token_ids']} for index {index} don't " | |
| f"match input ids {self.data[index]['input_ids']}", | |
| stacklevel=1, | |
| ) | |
| return SampleUnavailable( | |
| reason=f"Cached token ids do not match sample {index}", | |
| counts_as_failure=True, | |
| ) | |
| return { | |
| "hidden_states": loaded_hs["hidden_states"][:, :-1].flatten( | |
| 1 | |
| ), # [seq_len, 3 * hidden_size] | |
| "input_ids": loaded_hs["token_ids"], # [seq_len] | |
| "verifier_last_hidden_states": loaded_hs["hidden_states"][ | |
| :, -1 | |
| ], # [seq_len, hidden_size] | |
| "loss_mask": self.data[index]["loss_mask"], # [seq_len] | |
| } | |
| class CollateFn: | |
| """Picklable collate function for use with ``multiprocessing_context='spawn'``.""" | |
| def __init__( | |
| self, | |
| max_len: int, | |
| hidden_size: int, | |
| num_target_layers: int = 3, | |
| dtype: torch.dtype = torch.bfloat16, | |
| preprocess: Callable[[BatchType], BatchType] | None = None, | |
| ): | |
| self.max_len = max_len | |
| self.hidden_size = hidden_size | |
| self.num_target_layers = num_target_layers | |
| self.dtype = dtype | |
| self.preprocess = preprocess | |
| def _clean_batch( | |
| self, batch: Sequence[BatchType | SampleUnavailable | None] | |
| ) -> tuple[list[BatchType], list[SampleUnavailable], int]: | |
| """Preprocess valid samples and collect unavailable and dropped samples.""" | |
| preprocess = self.preprocess | |
| unavailable = [] | |
| num_dropped = 0 | |
| new_batch = [] | |
| for item in batch: | |
| if item is None: | |
| num_dropped += 1 | |
| continue | |
| if isinstance(item, SampleUnavailable): | |
| unavailable.append(item) | |
| num_dropped += 1 | |
| continue | |
| new_batch.append(preprocess(item) if preprocess else item) | |
| return new_batch, unavailable, num_dropped | |
| def __call__( | |
| self, batch: Sequence[BatchType | SampleUnavailable | None] | |
| ) -> BatchType: | |
| max_len = self.max_len | |
| dtype = self.dtype | |
| batch, unavailable, num_dropped = self._clean_batch(batch) | |
| if not batch: | |
| # Create empty sample which then gets padded to full | |
| # batch size if no valid samples are found. | |
| # Match the configured `dtype` so the placeholder doesn't crash | |
| # downstream layers loaded at a different precision (e.g. bf16 | |
| # weights vs fp32 default placeholders). | |
| empty = create_empty_sample( | |
| self.hidden_size, | |
| self.num_target_layers, | |
| dtype=dtype, | |
| ) | |
| if self.preprocess: | |
| empty = self.preprocess(empty) | |
| batch = [empty] | |
| locally_empty = True | |
| else: | |
| locally_empty = False | |
| collated_data: BatchType = {} | |
| for key in batch[0]: # type: ignore[union-attr] | |
| if key == "lengths": | |
| collated_data[key] = torch.cat([b[key] for b in batch], dim=0) # type: ignore[index] | |
| continue | |
| # one copy per sample: preallocated buffer, hidden states cast during write | |
| first = batch[0][key] # type: ignore[index] | |
| buffer_dtype = dtype if "hidden_states" in key else first.dtype | |
| out = torch.zeros( | |
| (max_len, *first.shape[1:]), dtype=buffer_dtype, device=first.device | |
| ) | |
| offset = 0 | |
| for b in batch: | |
| tensor = b[key] # type: ignore[index] | |
| num_rows = min(tensor.shape[0], max_len - offset) | |
| out[offset : offset + num_rows] = tensor[:num_rows] | |
| offset += num_rows | |
| if offset == max_len: | |
| break | |
| collated_data[key] = out.unsqueeze(0) | |
| # shape: [1, max_len, ...] | |
| # Include lengths until while they fit in max_len | |
| # The last included length is (if necessary) truncated | |
| # Any additional lengths are discarded | |
| lengths = collated_data.pop("lengths") | |
| new_lengths = [] | |
| cum_length = 0 | |
| for length in lengths: | |
| if length + cum_length >= max_len: | |
| new_lengths.append(max_len - cum_length) | |
| break | |
| new_lengths.append(length) | |
| cum_length += length | |
| lengths = torch.tensor(new_lengths, dtype=torch.long) | |
| # Create document_ids: maps each position to its document index, -1 for padding | |
| document_ids = torch.repeat_interleave( | |
| torch.arange(lengths.shape[0], dtype=torch.long), lengths | |
| ) | |
| document_ids = torch.cat( | |
| [ | |
| document_ids, | |
| -1 * torch.ones(max_len - document_ids.shape[0], dtype=torch.long), | |
| ] | |
| ).unsqueeze(0) | |
| # shape: [1, max_len] | |
| collated_data["document_ids"] = document_ids | |
| collated_data["error_records"] = num_dropped | |
| metadata = RecoveryMetadata.from_unavailable( | |
| unavailable, | |
| locally_empty=locally_empty, | |
| ) | |
| if metadata.failure_count or metadata.locally_empty: | |
| collated_data[RECOVERY_METADATA_KEY] = metadata | |
| return collated_data | |