Download source/src/speculators/data_generation/offline.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 4.13 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/data_generation/offline.py
- Command line
-
hf download hf://khazic/spec-b300/source/src/speculators/data_generation/offline.py
-
curl -L -o offline.py https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/data_generation/offline.py
4.13 kB
| import logging | |
| from pathlib import Path | |
| import torch | |
| logger = logging.getLogger(__name__) | |
| def check_hidden_states(data: dict, tokens: list[int]): | |
| required = {"token_ids", "hidden_states"} | |
| missing = required - data.keys() | |
| if missing: | |
| raise ValueError(f"Hidden-state payload is missing keys: {missing}") | |
| t_ids = data["token_ids"].tolist() | |
| if t_ids != tokens: | |
| raise ValueError(f"Token ids don't match expected token ids {tokens}") | |
| hs = data["hidden_states"] | |
| if not isinstance(hs, torch.Tensor): | |
| raise ValueError(f"Hidden states must be a tensor, got {type(hs).__name__}") | |
| if len(tokens) != hs.shape[0]: | |
| raise ValueError( | |
| f"Sequence length of hidden states {hs.shape[0]}" | |
| f" doesn't match num tokens {len(tokens)}" | |
| ) | |
| nan_count = 0 | |
| inf_count = 0 | |
| affected_layers: set[int] = set() | |
| rows_per_chunk = 256 | |
| for start in range(0, hs.shape[0], rows_per_chunk): | |
| # Process hidden states in chunks to avoid OOMs | |
| chunk = hs[start : start + rows_per_chunk] | |
| finite = torch.isfinite(chunk) | |
| if finite.all(): | |
| continue | |
| nan_count += int(torch.isnan(chunk).sum().item()) | |
| inf_count += int(torch.isinf(chunk).sum().item()) | |
| if hs.ndim >= 3: # noqa: PLR2004 | |
| bad_layers = (~finite).flatten(start_dim=2).any(dim=(0, 2)) | |
| affected_layers.update( | |
| bad_layers.nonzero(as_tuple=False).flatten().tolist() | |
| ) | |
| if nan_count or inf_count: | |
| details = ( | |
| f"shape={tuple(hs.shape)}, dtype={hs.dtype}, " | |
| f"nan_count={nan_count}, inf_count={inf_count}" | |
| ) | |
| if affected_layers: | |
| details += f", affected layer slots={sorted(affected_layers)}" | |
| raise ValueError(f"Hidden states contain non-finite values ({details})") | |
| def get_existing_hidden_state_indices(output_path: Path) -> list[int]: | |
| """Find existing `hs_i.safetensors` files (where i is the file index)""" | |
| existing_file_indices_set: set[int] = set() | |
| if not output_path.exists(): | |
| return [] | |
| for file_path in output_path.iterdir(): | |
| if file_path.name.startswith("hs_") and file_path.name.endswith(".safetensors"): | |
| index_str = file_path.stem[3:] # Remove "hs_" prefix | |
| try: | |
| file_index = int(index_str) | |
| existing_file_indices_set.add(file_index) | |
| except ValueError: | |
| continue | |
| return sorted(existing_file_indices_set) | |
| def get_indices_to_process( | |
| num_samples: int, | |
| max_samples: int | None, | |
| existing: list[int], | |
| world_size: int, | |
| rank: int, | |
| ) -> list[int]: | |
| """Determines which indices should be processed. If max_samples is None | |
| returns all dataset indices not in existing. Otherwise gets the first | |
| `max_samples - len(existing)` samples not already in existing. | |
| Args: | |
| num_samples: Total size of preprocessed dataset | |
| max_samples: (Optional) limit for number of samples to process | |
| existing: list of ids that have already been processed | |
| world_size: Number of nodes to generate on | |
| rank: The rank of the local node | |
| Returns: | |
| list of dataset indices to process | |
| """ | |
| target = min(max_samples, num_samples) if max_samples is not None else num_samples | |
| if target <= 0: | |
| return [] | |
| chunk_size = target // world_size | |
| remainder = target % world_size | |
| # Distribute remainder across the first `remainder` ranks so chunks differ | |
| # by at most 1. | |
| start = rank * chunk_size + min(rank, remainder) | |
| end = start + chunk_size + (1 if rank < remainder else 0) | |
| existing_s = set(existing) | |
| to_process = [i for i in range(start, end) if i not in existing_s] | |
| if not to_process: | |
| logger.info("All samples for this rank already processed!") | |
| return [] | |
| if len(existing_s & set(range(start, end))) > 0: | |
| logger.info( | |
| f"Found {len(existing_s & set(range(start, end)))} existing samples" | |
| f" for rank {rank}." | |
| ) | |
| return to_process | |