Upload scripts/fast4gpu_boost.py with huggingface_hub
Browse files- 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 |
-
|
| 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 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
)
|
| 120 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"))
|