khazic's picture
Archive three-epoch run: logs and provenance part 3
e65937c verified
Raw History Blame Contribute Delete
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