Spaces:
Sleeping
Sleeping
File size: 9,036 Bytes
c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 f2e829c c41af55 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 | """
fp8_cast_bf16_fallback.py
-------------------------
Standalone FP8 E4M3 โ BF16 dequantisation for DeepSeek-V3/Math-V2 style
safetensors checkpoints.
Processes ONE SHARD AT A TIME with explicit memory cleanup so that peak RAM
is bounded by the size of the largest single shard (~5-15 GB) rather than
the full model (~800 GB). Safe to run on standard HF Spaces hardware.
Usage:
python fp8_cast_bf16_fallback.py \
--input-fp8-hf-path /workspace/model-fp8 \
--output-bf16-hf-path /workspace/model-bf16
# Resume a partial run (already-converted shards are skipped):
python fp8_cast_bf16_fallback.py \
--input-fp8-hf-path /workspace/model-fp8 \
--output-bf16-hf-path /workspace/model-bf16 \
--resume
"""
import argparse
import gc
import json
import os
import shutil
from pathlib import Path
import torch
from safetensors import safe_open
from safetensors.torch import save_file
from tqdm import tqdm
# โโ Tile dequantisation โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def dequantize_fp8_tile(weight_fp8: torch.Tensor,
scale_inv: torch.Tensor) -> torch.Tensor:
"""
weight_fp8 : (rows, cols) in float8_e4m3fn
scale_inv : (ceil(rows/128), ceil(cols/128)) in float32
Returns : (rows, cols) in bfloat16
Processes in 128-row strips to avoid materialising a full (rows, cols)
float32 intermediate โ halves transient memory usage for large tensors.
"""
rows, cols = weight_fp8.shape
out = torch.empty((rows, cols), dtype=torch.bfloat16)
for row_start in range(0, rows, 128):
row_end = min(row_start + 128, rows)
tile_row = row_start // 128 # index into scale_inv rows
strip_fp8 = weight_fp8[row_start:row_end, :] # (โค128, cols)
# Build per-element scale for this strip
scale_strip = (
scale_inv[tile_row : tile_row + 1, :] # (1, tile_cols)
.repeat_interleave(128, dim=1)[:, :cols] # (1, cols)
.expand(row_end - row_start, -1) # (strip_rows, cols)
)
out[row_start:row_end, :] = (
strip_fp8.to(torch.float32).mul_(scale_strip).to(torch.bfloat16)
)
return out
# โโ Shard conversion โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def convert_shard(shard_path: Path,
out_path: Path,
fp8_key_names: set[str] | None = None) -> dict[str, str]:
"""
Convert a single safetensors shard in-place (one tensor at a time).
Returns a dict mapping tensor_name โ shard filename for index rebuilding.
"""
# โโ First pass: inventory the shard โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
key_list: list[str] = []
metadata: dict = {}
with safe_open(str(shard_path), framework="pt", device="cpu") as f:
key_list = list(f.keys())
try:
metadata = f.metadata() or {}
except Exception:
metadata = {}
# Detect FP8 keys in *this shard* (a key is FP8 if its _scale_inv is present)
all_keys_set = set(key_list)
fp8_keys_in_shard = {
k for k in key_list
if k.endswith(".weight")
and f"{k}_scale_inv" in all_keys_set
# Also accept an external hint (fp8_key_names) for cross-shard scales
or (fp8_key_names and k in fp8_key_names)
}
scale_keys = {f"{k}_scale_inv" for k in fp8_keys_in_shard}
# โโ Second pass: convert tensor by tensor โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
tensors_out: dict[str, torch.Tensor] = {}
with safe_open(str(shard_path), framework="pt", device="cpu") as f:
for key in key_list:
if key in scale_keys:
# Drop scale tensors from output
continue
tensor = f.get_tensor(key)
if key in fp8_keys_in_shard and tensor.dtype == torch.float8_e4m3fn:
# Load companion scale and dequantise
scale_inv = f.get_tensor(f"{key}_scale_inv").to(torch.float32)
tensor = dequantize_fp8_tile(tensor, scale_inv)
del scale_inv
tensors_out[key] = tensor
save_file(tensors_out, str(out_path), metadata=metadata)
weight_map = {k: shard_path.name for k in tensors_out}
# Explicit cleanup
del tensors_out
gc.collect()
return weight_map
# โโ Directory conversion โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def convert_directory(fp8_dir: Path, bf16_dir: Path, resume: bool = False) -> None:
bf16_dir.mkdir(parents=True, exist_ok=True)
# Collect shards
index_path = fp8_dir / "model.safetensors.index.json"
if index_path.exists():
with open(index_path) as f:
index = json.load(f)
shard_names = sorted(set(index["weight_map"].values()))
shard_paths = [fp8_dir / s for s in shard_names]
original_weight_map: dict[str, str] = index.get("weight_map", {})
else:
single = fp8_dir / "model.safetensors"
if not single.exists():
raise FileNotFoundError(f"No safetensors found in {fp8_dir}")
shard_paths = [single]
original_weight_map = {}
print(f"Found {len(shard_paths)} shard(s).")
print(f"Peak RAM per shard: ~{max(p.stat().st_size for p in shard_paths) * 2 / 1e9:.1f} GB "
f"(2ร largest shard on disk)")
combined_weight_map: dict[str, str] = {}
for shard_path in tqdm(shard_paths, desc="Converting shards", unit="shard"):
out_path = bf16_dir / shard_path.name
if resume and out_path.exists() and out_path.stat().st_size > 0:
tqdm.write(f" โท Skipping {shard_path.name} (already converted)")
# Still need to rebuild the weight map entry for the index
with safe_open(str(out_path), framework="pt", device="cpu") as f:
for key in f.keys():
combined_weight_map[key] = shard_path.name
continue
tqdm.write(f" โ {shard_path.name} ({shard_path.stat().st_size / 1e9:.1f} GB on disk)")
shard_map = convert_shard(shard_path, out_path)
combined_weight_map.update(shard_map)
tqdm.write(f" โ written โ {out_path.name} ({out_path.stat().st_size / 1e9:.1f} GB)")
# โโ Rebuild index โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
if index_path.exists():
# Exclude _scale_inv keys from the new weight map
new_weight_map = {
k: v for k, v in original_weight_map.items()
if not k.endswith("_scale_inv")
}
# Merge with actually-written map (handles edge cases)
new_weight_map.update(combined_weight_map)
new_index = {
"metadata": index.get("metadata", {}),
"weight_map": new_weight_map,
}
with open(bf16_dir / "model.safetensors.index.json", "w") as f:
json.dump(new_index, f, indent=2)
print("Index written.")
# โโ Copy config / tokeniser files โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
for pattern in [
"config.json",
"generation_config.json",
"tokenizer*.json",
"tokenizer*.model",
"special_tokens_map.json",
"*.py",
]:
for src in fp8_dir.glob(pattern):
dst = bf16_dir / src.name
if not dst.exists():
shutil.copy2(src, dst)
print(f"Copied {src.name}")
print(f"\nโ BF16 model written to: {bf16_dir}")
# โโ CLI โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def main() -> None:
parser = argparse.ArgumentParser(
description="Dequantise DeepSeek FP8 safetensors โ BF16 (low RAM, shard-by-shard)"
)
parser.add_argument("--input-fp8-hf-path", required=True, type=Path)
parser.add_argument("--output-bf16-hf-path", required=True, type=Path)
parser.add_argument(
"--resume",
action="store_true",
help="Skip shards whose output file already exists (useful after a crash)",
)
args = parser.parse_args()
convert_directory(args.input_fp8_hf_path, args.output_bf16_hf_path, args.resume)
if __name__ == "__main__":
main()
|