[CI/Build] Consolidate model loader tests and requirements (#25765)
Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
@@ -639,6 +639,19 @@ def runai_safetensors_weights_iterator(
|
||||
yield from tensor_iter
|
||||
|
||||
|
||||
def _init_loader(
|
||||
pg: torch.distributed.ProcessGroup,
|
||||
device: torch.device,
|
||||
f_list: list[str],
|
||||
*,
|
||||
nogds: bool = False,
|
||||
):
|
||||
loader = SafeTensorsFileLoader(pg, device, nogds=nogds)
|
||||
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
|
||||
loader.add_filenames(rank_file_map)
|
||||
return loader
|
||||
|
||||
|
||||
def fastsafetensors_weights_iterator(
|
||||
hf_weights_files: list[str],
|
||||
use_tqdm_on_load: bool,
|
||||
@@ -656,17 +669,31 @@ def fastsafetensors_weights_iterator(
|
||||
for i in range(0, len(hf_weights_files), pg.size())
|
||||
]
|
||||
|
||||
nogds = False
|
||||
|
||||
for f_list in tqdm(
|
||||
weight_files_sub_lists,
|
||||
desc="Loading safetensors using Fastsafetensor loader",
|
||||
disable=not enable_tqdm(use_tqdm_on_load),
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
loader = SafeTensorsFileLoader(pg, device)
|
||||
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
|
||||
loader.add_filenames(rank_file_map)
|
||||
loader = _init_loader(pg, device, f_list, nogds=nogds)
|
||||
try:
|
||||
fb = loader.copy_files_to_device()
|
||||
try:
|
||||
fb = loader.copy_files_to_device()
|
||||
except RuntimeError as e:
|
||||
if "gds" not in str(e):
|
||||
raise
|
||||
|
||||
loader.close()
|
||||
nogds = True
|
||||
logger.warning_once(
|
||||
"GDS not enabled, setting `nogds=True`.\n"
|
||||
"For more information, see: https://github.com/foundation-model-stack/fastsafetensors?tab=readme-ov-file#basic-api-usages"
|
||||
)
|
||||
loader = _init_loader(pg, device, f_list, nogds=nogds)
|
||||
fb = loader.copy_files_to_device()
|
||||
|
||||
try:
|
||||
keys = list(fb.key_to_rank_lidx.keys())
|
||||
for k in keys:
|
||||
|
||||
Reference in New Issue
Block a user