Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 11 additions & 6 deletions inference/fp8_cast_bf16.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import os
import json
from argparse import ArgumentParser
from collections import OrderedDict
from glob import glob
from tqdm import tqdm

Expand Down Expand Up @@ -36,8 +37,8 @@ def main(fp8_path, bf16_path):
model_index = json.load(f)
weight_map = model_index["weight_map"]

# Cache for loaded safetensor files
loaded_files = {}
# Cache for loaded safetensor files (OrderedDict for LRU eviction)
loaded_files = OrderedDict()
fp8_weight_names = []

# Helper function to get tensor from the correct file
Expand All @@ -58,6 +59,9 @@ def get_tensor(tensor_name):
if file_name not in loaded_files:
file_path = os.path.join(fp8_path, file_name)
loaded_files[file_name] = load_file(file_path, device="cuda")
else:
# Mark as recently used for LRU eviction
loaded_files.move_to_end(file_name)
return loaded_files[file_name][tensor_name]

safetensor_files = list(glob(os.path.join(fp8_path, "*.safetensors")))
Expand All @@ -66,6 +70,7 @@ def get_tensor(tensor_name):
file_name = os.path.basename(safetensor_file)
current_state_dict = load_file(safetensor_file, device="cuda")
loaded_files[file_name] = current_state_dict
loaded_files.move_to_end(file_name)

new_state_dict = {}
for weight_name, weight in current_state_dict.items():
Expand All @@ -87,10 +92,10 @@ def get_tensor(tensor_name):
new_safetensor_file = os.path.join(bf16_path, file_name)
save_file(new_state_dict, new_safetensor_file)

# Memory management: keep only the 2 most recently used files
if len(loaded_files) > 2:
oldest_file = next(iter(loaded_files))
del loaded_files[oldest_file]
# Memory management: evict least recently used files, keep at most 2
while len(loaded_files) > 2:
# Evict the least recently used entry (first item in OrderedDict)
loaded_files.popitem(last=False)
torch.cuda.empty_cache()

# Update model index
Expand Down