Diffusers
Safetensors
HY / predictor_data /writer.py
Cccccz's picture
Upload batch 64: 500 files (0.40 GiB)
5f0e4a2 verified
Raw History Blame Contribute Delete
4.38 kB
"""Atomic safetensors writer and JSONL manifest for Predictor data."""
from __future__ import annotations
import json
import os
from pathlib import Path
from typing import Any, Mapping
import torch
from safetensors.torch import save_file
from .schema import SCHEMA_VERSION, validate_case_tensors, validate_chunk_tensors
def _cpu_contiguous(tensor: torch.Tensor, *, dtype: torch.dtype | None = None) -> torch.Tensor:
tensor = tensor.detach()
if dtype is not None and tensor.is_floating_point():
tensor = tensor.to(dtype=dtype)
return tensor.to(device="cpu").contiguous()
class PredictorDatasetWriter:
def __init__(self, root: str | os.PathLike[str], *, save_dtype: torch.dtype = torch.bfloat16):
self.root = Path(root).resolve()
self.case_dir = self.root / "cases"
self.chunk_dir = self.root / "chunks"
self.manifest_path = self.root / "manifest.jsonl"
self.train_eval_manifest_path = self.root / "train_eval_manifest.jsonl"
self.save_dtype = save_dtype
self.case_dir.mkdir(parents=True, exist_ok=True)
self.chunk_dir.mkdir(parents=True, exist_ok=True)
self._manifest_keys = self._load_existing_keys()
def _load_existing_keys(self) -> set[tuple[int, int, int]]:
keys: set[tuple[int, int, int]] = set()
if not self.manifest_path.exists():
return keys
with self.manifest_path.open("r", encoding="utf-8") as handle:
for line in handle:
if not line.strip():
continue
item = json.loads(line)
keys.add((int(item["case_id"]), int(item["seed"]), int(item["chunk_id"])))
return keys
def case_path(self, case_id: int) -> Path:
return self.case_dir / f"case_{case_id:02d}.safetensors"
def chunk_path(self, case_id: int, seed: int, chunk_id: int) -> Path:
return self.chunk_dir / f"case_{case_id:02d}_seed_{seed}" / f"chunk_{chunk_id:02d}.safetensors"
def is_chunk_complete(self, case_id: int, seed: int, chunk_id: int) -> bool:
key = (case_id, seed, chunk_id)
return key in self._manifest_keys and self.chunk_path(case_id, seed, chunk_id).is_file()
def _atomic_save(self, path: Path, tensors: Mapping[str, torch.Tensor]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp_path = path.with_suffix(path.suffix + f".tmp.{os.getpid()}")
save_file(dict(tensors), str(tmp_path))
os.replace(tmp_path, path)
def save_case(self, case_id: int, tensors: Mapping[str, torch.Tensor]) -> Path:
converted = {
name: _cpu_contiguous(tensor, dtype=self.save_dtype)
for name, tensor in tensors.items()
}
validate_case_tensors(converted)
path = self.case_path(case_id)
self._atomic_save(path, converted)
return path
def save_chunk(
self,
*,
case_id: int,
seed: int,
chunk_id: int,
tensors: Mapping[str, torch.Tensor],
metadata: Mapping[str, Any],
) -> Path:
key = (case_id, seed, chunk_id)
if self.is_chunk_complete(*key):
return self.chunk_path(*key)
converted = {}
for name, tensor in tensors.items():
dtype = self.save_dtype if tensor.is_floating_point() and "timestep" not in name else None
if "timestep" in name:
dtype = torch.float32
converted[name] = _cpu_contiguous(tensor, dtype=dtype)
validate_chunk_tensors(converted)
path = self.chunk_path(*key)
self._atomic_save(path, converted)
record = {
"schema_version": SCHEMA_VERSION,
"case_id": case_id,
"seed": seed,
"chunk_id": chunk_id,
"tensor_file": str(path.relative_to(self.root)),
"case_tensor_file": str(self.case_path(case_id).relative_to(self.root)),
**dict(metadata),
}
line = json.dumps(record, ensure_ascii=False, sort_keys=True) + "\n"
for manifest in (self.manifest_path, self.train_eval_manifest_path):
with manifest.open("a", encoding="utf-8") as handle:
handle.write(line)
handle.flush()
os.fsync(handle.fileno())
self._manifest_keys.add(key)
return path