thefinalboss commited on
Commit
9f2620b
·
verified ·
1 Parent(s): bac4ce7

Upload scripts/fast4gpu_boost.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. scripts/fast4gpu_boost.py +17 -6
scripts/fast4gpu_boost.py CHANGED
@@ -105,7 +105,7 @@ eng = eng.to(device)
105
  eng.reset_thought(B)
106
 
107
  try:
108
- eng = torch.compile(eng, mode="reduce-overhead")
109
  print(f"GPU {GPU}: compile OK", flush=True)
110
  except Exception as e:
111
  print(f"GPU {GPU}: compile skip: {e}", flush=True)
@@ -113,11 +113,22 @@ except Exception as e:
113
  opt = torch.optim.SGD(eng.parameters(), lr=LR, momentum=0.9)
114
 
115
  if not SHARD.exists():
116
- raise FileNotFoundError(
117
- f"Shard not found: {SHARD}\n"
118
- "Place tokenized int64 1D shard at data/shard_gpu{id}.pt or set SHARD="
119
- )
120
- tokens = torch.load(SHARD, weights_only=False).to(torch.int64)
 
 
 
 
 
 
 
 
 
 
 
121
  step_tokens = B * SEQ
122
 
123
  start_token = int(os.environ.get("START_TOKEN", "0"))
 
105
  eng.reset_thought(B)
106
 
107
  try:
108
+ print("compile disabled for VRAM", flush=True)
109
  print(f"GPU {GPU}: compile OK", flush=True)
110
  except Exception as e:
111
  print(f"GPU {GPU}: compile skip: {e}", flush=True)
 
113
  opt = torch.optim.SGD(eng.parameters(), lr=LR, momentum=0.9)
114
 
115
  if not SHARD.exists():
116
+ # allow .npy memmap fallback for large phase2 shards
117
+ alt = Path(str(SHARD) + ".npy") if not str(SHARD).endswith(".npy") else SHARD
118
+ npy = SHARD if str(SHARD).endswith(".npy") else Path(str(SHARD).replace(".pt", ".npy"))
119
+ if not npy.exists():
120
+ raise FileNotFoundError(
121
+ f"Shard not found: {SHARD} (also tried {npy})"
122
+ )
123
+ SHARD = npy
124
+
125
+ if str(SHARD).endswith(".npy"):
126
+ import numpy as np
127
+ tokens = torch.from_numpy(np.load(str(SHARD), mmap_mode="r")).to(torch.int64)
128
+ print(f"GPU {GPU}: memmap shard {SHARD} len={len(tokens):,}", flush=True)
129
+ else:
130
+ tokens = torch.load(SHARD, weights_only=False).to(torch.int64)
131
+ print(f"GPU {GPU}: loaded shard {SHARD} len={len(tokens):,}", flush=True)
132
  step_tokens = B * SEQ
133
 
134
  start_token = int(os.environ.get("START_TOKEN", "0"))