Spaces:
Running on Zero
Running on Zero
Table to RAM with 64 MB sequential reads (page-faulting the mmap over the mount was far too slow)
Browse files- app.py +9 -3
- whittle_load.py +22 -1
app.py
CHANGED
|
@@ -98,9 +98,15 @@ def _table_to_ram():
|
|
| 98 |
if not isinstance(PLE.ngram_embedding, MmapNGramTable):
|
| 99 |
return
|
| 100 |
t0 = time.time()
|
| 101 |
-
table
|
| 102 |
-
|
| 103 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
|
| 105 |
|
| 106 |
threading.Thread(target=_table_to_ram, daemon=True).start()
|
|
|
|
| 98 |
if not isinstance(PLE.ngram_embedding, MmapNGramTable):
|
| 99 |
return
|
| 100 |
t0 = time.time()
|
| 101 |
+
print("[table] reading the n-gram table into RAM (64 MB sequential reads)", flush=True)
|
| 102 |
+
try:
|
| 103 |
+
emb = PLE.ngram_embedding.to_ram()
|
| 104 |
+
except Exception as e:
|
| 105 |
+
print(f"[table] staying memory-mapped: {e!r}", flush=True)
|
| 106 |
+
return
|
| 107 |
+
PLE.ngram_embedding = emb
|
| 108 |
+
dt = time.time() - t0
|
| 109 |
+
print(f"[table] n-gram table in RAM: {tuple(emb.weight.shape)} in {dt:.0f}s ({emb.weight.numel() * emb.weight.element_size() / 2**30 / dt:.2f} GiB/s)", flush=True)
|
| 110 |
|
| 111 |
|
| 112 |
threading.Thread(target=_table_to_ram, daemon=True).start()
|
whittle_load.py
CHANGED
|
@@ -55,7 +55,7 @@ class MmapNGramTable(torch.nn.Module):
|
|
| 55 |
|
| 56 |
def __init__(self, pieces, dim: int, dtype):
|
| 57 |
super().__init__()
|
| 58 |
-
self._parts, starts, row = [], [], 0
|
| 59 |
for path, offset, rows, stored in pieces:
|
| 60 |
with open(path, "rb") as fh:
|
| 61 |
buf = mmap.mmap(fh.fileno(), 0, access=mmap.ACCESS_READ)
|
|
@@ -65,6 +65,27 @@ class MmapNGramTable(torch.nn.Module):
|
|
| 65 |
self.register_buffer("_bounds", torch.tensor(starts[1:], dtype=torch.long), persistent=False)
|
| 66 |
self.register_buffer("weight", torch.empty(0, dim, dtype=dtype), persistent=False) # device/dtype marker only
|
| 67 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
| 69 |
flat = ids.reshape(-1).cpu()
|
| 70 |
which = torch.bucketize(flat, self._bounds, right=True)
|
|
|
|
| 55 |
|
| 56 |
def __init__(self, pieces, dim: int, dtype):
|
| 57 |
super().__init__()
|
| 58 |
+
self._pieces, self._parts, starts, row = list(pieces), [], [], 0
|
| 59 |
for path, offset, rows, stored in pieces:
|
| 60 |
with open(path, "rb") as fh:
|
| 61 |
buf = mmap.mmap(fh.fileno(), 0, access=mmap.ACCESS_READ)
|
|
|
|
| 65 |
self.register_buffer("_bounds", torch.tensor(starts[1:], dtype=torch.long), persistent=False)
|
| 66 |
self.register_buffer("weight", torch.empty(0, dim, dtype=dtype), persistent=False) # device/dtype marker only
|
| 67 |
|
| 68 |
+
def to_ram(self, chunk: int = 64 << 20) -> torch.nn.Embedding:
|
| 69 |
+
"""Copy the whole table into one in-memory nn.Embedding with large sequential file reads (much faster than
|
| 70 |
+
faulting the memory map in page by page, especially on a network mount)."""
|
| 71 |
+
stored = self._pieces[0][3]
|
| 72 |
+
assert all(p[3] == stored for p in self._pieces), "pieces with mixed dtypes"
|
| 73 |
+
table = torch.empty((self.num_embeddings, self.embedding_dim), dtype=stored)
|
| 74 |
+
raw = memoryview(table.view(torch.uint8).numpy().reshape(-1))
|
| 75 |
+
pos = 0
|
| 76 |
+
for path, offset, rows, _ in self._pieces:
|
| 77 |
+
nbytes = rows * self.embedding_dim * table.element_size()
|
| 78 |
+
with open(path, "rb", buffering=0) as fh:
|
| 79 |
+
fh.seek(offset)
|
| 80 |
+
done = 0
|
| 81 |
+
while done < nbytes:
|
| 82 |
+
n = fh.readinto(raw[pos + done: pos + min(nbytes, done + chunk)])
|
| 83 |
+
if not n:
|
| 84 |
+
raise IOError(f"short read in {path} at {offset + done}")
|
| 85 |
+
done += n
|
| 86 |
+
pos += nbytes
|
| 87 |
+
return torch.nn.Embedding.from_pretrained(table.to(self.weight.dtype), freeze=True)
|
| 88 |
+
|
| 89 |
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
| 90 |
flat = ids.reshape(-1).cpu()
|
| 91 |
which = torch.bucketize(flat, self._bounds, right=True)
|