Download training/tensor_cache.py from gaaaaaaaaaaa/MultimodalReasoning3B: direct link, hf CLI and curl.
- Browser
- Download file 5.24 kB
-
https://huggingface.co/gaaaaaaaaaaa/MultimodalReasoning3B/resolve/main/training/tensor_cache.py
- Command line
-
hf download hf://gaaaaaaaaaaa/MultimodalReasoning3B/training/tensor_cache.py
-
curl -L -o tensor_cache.py https://huggingface.co/gaaaaaaaaaaa/MultimodalReasoning3B/resolve/main/training/tensor_cache.py
5.24 kB
| from __future__ import annotations | |
| import hashlib | |
| import importlib.util | |
| import json | |
| from pathlib import Path | |
| from typing import Any, Iterable, Sequence | |
| CACHE_VERSION = 1 | |
| def _json_hash(value: Any) -> str: | |
| encoded = json.dumps(value, ensure_ascii=False, sort_keys=True).encode("utf-8") | |
| return hashlib.sha256(encoded).hexdigest() | |
| def _hash_text(value: Any) -> str: | |
| return hashlib.sha256(str(value or "").encode("utf-8")).hexdigest() | |
| def parquet_reader_available() -> bool: | |
| return any( | |
| importlib.util.find_spec(package) is not None | |
| for package in ("pyarrow", "polars", "pandas") | |
| ) | |
| def dataset_source_files(dataset_name: str, data_dir: str | Path) -> list[Path]: | |
| root = Path(data_dir) | |
| if dataset_name == "multimodal": | |
| return sorted((root / "MultimodalReasoning").glob("*.jsonl")) | |
| if dataset_name == "reasoning_dpo": | |
| return sorted((root / "ReasoningDPO").glob("*.jsonl")) | |
| if dataset_name == "limo": | |
| return [root / "limo.jsonl"] | |
| if dataset_name == "complex_bespoke": | |
| return [root / "ComplexReasoningBespoke.parquet"] | |
| if dataset_name == "all": | |
| files: list[Path] = [] | |
| files.extend(dataset_source_files("multimodal", root)) | |
| files.extend(dataset_source_files("reasoning_dpo", root)) | |
| files.extend(dataset_source_files("limo", root)) | |
| files.extend(dataset_source_files("complex_bespoke", root)) | |
| return files | |
| return [] | |
| def file_metadata(path: Path) -> dict[str, Any]: | |
| if not path.exists(): | |
| return { | |
| "path": str(path), | |
| "exists": False, | |
| } | |
| stat = path.stat() | |
| return { | |
| "path": str(path), | |
| "exists": True, | |
| "size": stat.st_size, | |
| "mtime_ns": stat.st_mtime_ns, | |
| } | |
| def data_digest(records: Sequence[dict[str, str]]) -> dict[str, Any]: | |
| digest = hashlib.sha256() | |
| for record in records: | |
| digest.update(str(record.get("user", "")).encode("utf-8")) | |
| digest.update(b"\0") | |
| digest.update(str(record.get("assistant", "")).encode("utf-8")) | |
| digest.update(b"\0\0") | |
| return { | |
| "rows": len(records), | |
| "sha256": digest.hexdigest(), | |
| } | |
| def tokenizer_metadata(tokenizer: Any) -> dict[str, Any]: | |
| return { | |
| "name_or_path": str(getattr(tokenizer, "name_or_path", "")), | |
| "class": tokenizer.__class__.__name__, | |
| "eos_token": str(getattr(tokenizer, "eos_token", "")), | |
| "pad_token": str(getattr(tokenizer, "pad_token", "")), | |
| "padding_side": str(getattr(tokenizer, "padding_side", "")), | |
| "chat_template_sha256": _hash_text(getattr(tokenizer, "chat_template", "")), | |
| "vocab_size": len(tokenizer) if hasattr(tokenizer, "__len__") else None, | |
| } | |
| def build_cache_metadata( | |
| *, | |
| config: dict[str, Any], | |
| tokenizer: Any, | |
| provided_records: Sequence[dict[str, str]] | None = None, | |
| ) -> dict[str, Any]: | |
| relevant_config = { | |
| key: config.get(key) | |
| for key in ( | |
| "model_name", | |
| "data_dir", | |
| "dataset_name", | |
| "parquet_limit", | |
| "system_prompt", | |
| "dataset_format", | |
| "append_eos_to_completion", | |
| "max_samples", | |
| "eval_split_size", | |
| "max_eval_samples", | |
| "shuffle_data", | |
| "seed", | |
| "max_length", | |
| "completion_only_loss", | |
| "assistant_only_loss", | |
| ) | |
| } | |
| metadata: dict[str, Any] = { | |
| "cache_version": CACHE_VERSION, | |
| "config": relevant_config, | |
| "tokenizer": tokenizer_metadata(tokenizer), | |
| "parquet_reader_available": parquet_reader_available(), | |
| } | |
| if provided_records is not None: | |
| metadata["provided_data"] = data_digest(provided_records) | |
| else: | |
| metadata["source_files"] = [ | |
| file_metadata(path) | |
| for path in dataset_source_files( | |
| str(config.get("dataset_name", "all")), | |
| str(config.get("data_dir", "data")), | |
| ) | |
| ] | |
| metadata["cache_key"] = _json_hash(metadata)[:24] | |
| return metadata | |
| def cache_path(cache_root: str | Path, metadata: dict[str, Any]) -> Path: | |
| dataset_name = str(metadata["config"].get("dataset_name") or "dataset") | |
| safe_name = "".join( | |
| char if char.isalnum() or char in ("-", "_") else "_" | |
| for char in dataset_name | |
| ) | |
| return Path(cache_root) / f"{safe_name}_{metadata['cache_key']}" | |
| def metadata_path(path: str | Path) -> Path: | |
| return Path(path) / "metadata.json" | |
| def read_metadata(path: str | Path) -> dict[str, Any] | None: | |
| meta_path = metadata_path(path) | |
| if not meta_path.exists(): | |
| return None | |
| with meta_path.open("r", encoding="utf-8") as f: | |
| return json.load(f) | |
| def write_metadata(path: str | Path, metadata: dict[str, Any]) -> None: | |
| meta_path = metadata_path(path) | |
| meta_path.parent.mkdir(parents=True, exist_ok=True) | |
| with meta_path.open("w", encoding="utf-8") as f: | |
| json.dump(metadata, f, ensure_ascii=False, indent=2, sort_keys=True) | |
| def metadata_matches(path: str | Path, metadata: dict[str, Any]) -> bool: | |
| cached = read_metadata(path) | |
| return cached == metadata | |