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