LazyChunkedTensor

This commit is contained in:
Xuan Son Nguyen
2026-08-27 18:29:28 +02:00
parent fbc5135abb
commit a510c82e17
3 changed files with 98 additions and 84 deletions
+6 -2
View File
@@ -1006,12 +1006,16 @@ class ModelBase:
else:
raise ValueError(f"Unknown file type: {self.ftype.name}")
# a chunked tensor quantizes as one chunk at a time, while it is written
quantize = data.quantize if isinstance(data, gguf.LazyChunkedTensor) else (
lambda qtype, d=data: gguf.quants.quantize(d, qtype))
try:
data = gguf.quants.quantize(data, data_qtype)
data = quantize(data_qtype)
except gguf.QuantError as e:
logger.warning("%s, %s", e, "falling back to F16")
data_qtype = gguf.GGMLQuantizationType.F16
data = gguf.quants.quantize(data, data_qtype)
data = quantize(data_qtype)
shape = gguf.quant_shape_from_byte_shape(data.shape, data_qtype) if data.dtype == np.uint8 else data.shape
+31 -82
View File
@@ -1,6 +1,5 @@
from __future__ import annotations
from pathlib import Path
from typing import Iterable
import torch
@@ -32,13 +31,9 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# shards held only until the row stride is known, normally none
self._ple_pending: dict[int, Tensor] = {}
self._ple_shard_rows: dict[int, int] = {}
# only the shard names, so the table itself is never held
self._ple_shards: dict[int, str] = {}
self._ple_row_dim: int | None = None
self._ple_rows_per_shard: int | None = None
self._ple_map: np.memmap | None = None
self._ple_path: Path | None = None
def _read_hash_constants(self, suffix: str) -> list[int]:
"""Read an int64 PLE constant straight from the checkpoint.
@@ -145,100 +140,54 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
return super().modify_tensors(data_torch, name, bid)
# -- the PLE table ----------------------------------------------------
#
# The 128 shards concatenate into one enormous tensor, which peaks near 300 GB of RSS.
# Each shard is written straight into a memory-mapped file at its final row offset and
# then dropped, so only one shard is resident. The file is removed after the write.
# It holds float32 because base.py has already cast the shards to it.
# the shards concatenate into a tensor of well over 100 GB
# use LazyChunkedTensor here, a single shard resident at a time
def _place_ple_shard(self, data_torch: Tensor, name: str) -> Iterable[tuple[str, Tensor]]:
idx = int(name.rpartition(".shard_")[2].partition(".")[0])
n_parts = self.hparams["split_ngram_parts"]
rows, row_dim = int(data_torch.shape[0]), int(data_torch.shape[-1])
self._ple_row_dim = row_dim
self._ple_shard_rows[idx] = rows
self._ple_shards[idx] = name
self._ple_row_dim = int(data_torch.shape[-1])
if self._ple_map is None:
if idx == n_parts - 1 and n_parts > 1:
# the last shard can be short, so it cannot set the stride
# this happens only if the checkpoint yields the shards out of order
self._ple_pending[idx] = data_torch
return []
self._ple_rows_per_shard = rows
self._ple_path = self.fname_out.parent / f".{self.fname_out.stem}.ple.tmp"
self._ple_map = np.memmap(
self._ple_path, dtype=np.float32, mode="w+",
shape=(n_parts * rows, row_dim))
for i, held in list(self._ple_pending.items()):
self._ple_pending.pop(i)
self._write_ple_shard(i, held)
self._write_ple_shard(idx, data_torch)
if len(self._ple_shard_rows) < n_parts:
if len(self._ple_shards) < n_parts:
return []
total = sum(self._ple_shard_rows.values())
table = self._finish_ple_table(total)
# the checkpoint may yield the shards in any order, the row order is by index
shards = [self._ple_shards[i] for i in sorted(self._ple_shards)]
rows = 0
for shard in shards:
shape = self.model_tensors[shard]().shape
if int(shape[-1]) != self._ple_row_dim:
raise ValueError(
f"PLE shard {shard} has row dim {int(shape[-1])}, expected {self._ple_row_dim}")
rows += int(shape[0])
table = gguf.LazyChunkedTensor(
[self._load_ple_shard(shard) for shard in shards],
shape=(rows, self._ple_row_dim),
dtype=np.float32,
)
gguf_name = gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.PER_LAYER_TOKEN_EMBD]
return [(gguf_name + ".weight", table)]
def _write_ple_shard(self, idx: int, shard: Tensor) -> None:
# the caller opens the map and fixes the stride before the first write
assert self._ple_map is not None and self._ple_rows_per_shard is not None
def _load_ple_shard(self, name: str):
def load() -> np.ndarray:
from .base import LazyTorchTensor
rows = int(shard.shape[0])
if idx != self.hparams["split_ngram_parts"] - 1 and rows != self._ple_rows_per_shard:
raise ValueError(
f"PLE shard {idx} has {rows} rows, expected {self._ple_rows_per_shard}; "
"shards other than the last must be uniform for direct placement"
)
start = idx * self._ple_rows_per_shard
# the shard is still lazy here; force it, so exactly one shard is resident
from .base import LazyTorchTensor
eager = LazyTorchTensor.to_eager(shard).to(torch.float32).contiguous()
self._ple_map[start:start + rows] = eager.numpy()
del eager
def _finish_ple_table(self, total_rows: int):
# only reached once every shard has been written, so the map is open
assert self._ple_map is not None and self._ple_path is not None
assert self._ple_row_dim is not None
self._ple_map.flush()
del self._ple_map
self._ple_map = None
# trim the tail if the last shard came up short of a full stride
want = total_rows * self._ple_row_dim * 4
if self._ple_path.stat().st_size != want:
with open(self._ple_path, "r+b") as f:
f.truncate(want)
raw = np.memmap(self._ple_path, dtype=np.float32, mode="r+",
shape=(total_rows, self._ple_row_dim))
return torch.from_numpy(np.asarray(raw))
# a fresh lazy tensor every call, or to_eager() memoizes every shard
eager = LazyTorchTensor.to_eager(self.model_tensors[name]())
return eager.to(torch.float32).contiguous().numpy()
return load
def prepare_tensors(self):
super().prepare_tensors()
if self._ple_pending:
n_parts = self.hparams.get("split_ngram_parts", 0)
if self._ple_shards and len(self._ple_shards) != n_parts:
raise ValueError(
f"unprocessed PLE embedding shards: {sorted(self._ple_pending)}"
f"got {len(self._ple_shards)} PLE embedding shards, expected {n_parts}"
)
def write(self):
try:
super().write()
finally:
if self._ple_path is not None and self._ple_path.exists():
self._ple_path.unlink()
@ModelBase.register("Qwen4ExpForConditionalGeneration")
@ModelBase.example("Qwen/Qwen3.8-Flash-Next")
+61
View File
@@ -226,3 +226,64 @@ class LazyNumpyTensor(LazyBase):
return eager.tofile(*args, **kwargs)
# TODO: __array_function__
# Tensor written to file one row-chunk at a time
class LazyChunkedTensor:
def __init__(
self, chunks: list[Callable[[], np.ndarray]], shape: tuple[int, ...], dtype: DTypeLike,
qtype: Any = None, byteswap: bool = False,
):
self._chunks = chunks
self._qtype = qtype
self._byteswap = byteswap
self.shape = tuple(shape)
self.dtype = np.dtype(dtype)
@property
def nbytes(self) -> int:
n = self.dtype.itemsize
for d in self.shape:
n *= d
return n
def numpy(self) -> LazyChunkedTensor:
return self
def quantize(self, qtype: Any) -> LazyChunkedTensor:
from .constants import GGMLQuantizationType
from .quants import QuantError, quant_shape_to_byte_shape
if qtype == GGMLQuantizationType.F32:
shape, dtype = self.shape, np.dtype(np.float32)
elif qtype == GGMLQuantizationType.F16:
shape, dtype = self.shape, np.dtype(np.float16)
else:
try:
shape, dtype = quant_shape_to_byte_shape(self.shape, qtype), np.dtype(np.uint8)
except ValueError as e:
# raised here and not per chunk, so callers can still fall back to F16
raise QuantError(str(e)) from e
return LazyChunkedTensor(self._chunks, shape, dtype, qtype, self._byteswap)
def byteswap(self, inplace: bool = False) -> LazyChunkedTensor:
if inplace:
raise NotImplementedError("a chunked tensor cannot be byteswapped in place")
return LazyChunkedTensor(self._chunks, self.shape, self.dtype, self._qtype, not self._byteswap)
def tofile(self, *args, **kwargs) -> None:
from .quants import quantize
written = 0
for load_chunk in self._chunks:
chunk = load_chunk()
if self._qtype is not None:
# exact only because chunks split on rows, and blocks never cross one
chunk = quantize(chunk, self._qtype)
if self._byteswap:
chunk = chunk.byteswap(inplace=False)
chunk.tofile(*args, **kwargs)
written += chunk.nbytes
del chunk
assert written == self.nbytes, f"chunked tensor wrote {written} bytes, expected {self.nbytes}"