logic65 commited on
Commit
e846b6b
·
verified ·
1 Parent(s): 8792537

Table to RAM with 64 MB sequential reads (page-faulting the mmap over the mount was far too slow)

Browse files
Files changed (2) hide show
  1. app.py +9 -3
  2. 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 = torch.cat(PLE.ngram_embedding._parts) # sequential read of the pieces from the mount
102
- PLE.ngram_embedding = torch.nn.Embedding.from_pretrained(table, freeze=True)
103
- print(f"[table] n-gram table in RAM: {tuple(table.shape)} in {time.time() - t0:.0f}s", flush=True)
 
 
 
 
 
 
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)