Fix hf push

This commit is contained in:
litagin02
2024-02-17 11:31:35 +09:00
parent ff730476e7
commit 8b76109b75
2 changed files with 43 additions and 10 deletions

View File

@@ -6,6 +6,7 @@ import platform
import torch import torch
import torch.distributed as dist import torch.distributed as dist
from huggingface_hub import HfApi
from torch.cuda.amp import GradScaler, autocast from torch.cuda.amp import GradScaler, autocast
from torch.nn import functional as F from torch.nn import functional as F
from torch.nn.parallel import DistributedDataParallel as DDP from torch.nn.parallel import DistributedDataParallel as DDP
@@ -45,6 +46,8 @@ torch.backends.cuda.enable_math_sdp(True)
global_step = 0 global_step = 0
api = HfApi()
def run(): def run():
# Command line configuration is not recommended unless necessary, use config.yml # Command line configuration is not recommended unless necessary, use config.yml
@@ -483,6 +486,18 @@ def run():
), ),
for_infer=True, for_infer=True,
) )
if hps.repo_id is not None:
api.upload_folder(
repo_id=hps.repo_id,
folder_path=model_dir,
path_in_repo=f"Data/{config.model_name}/models",
delete_patterns="*.pth",
)
api.upload_folder(
repo_id=hps.repo_id,
folder_path=config.out_dir,
path_in_repo=f"model_assets/{config.model_name}",
)
if pbar is not None: if pbar is not None:
pbar.close() pbar.close()
@@ -765,6 +780,18 @@ def train_and_evaluate(
), ),
for_infer=True, for_infer=True,
) )
if hps.repo_id is not None:
api.upload_folder(
repo_id=hps.repo_id,
folder_path=hps.model_dir,
path_in_repo=f"Data/{config.model_name}/models",
delete_patterns="*.pth",
)
api.upload_folder(
repo_id=hps.repo_id,
folder_path=config.out_dir,
path_in_repo=f"model_assets/{config.model_name}",
)
global_step += 1 global_step += 1
if pbar is not None: if pbar is not None:

View File

@@ -578,10 +578,13 @@ def run():
api.upload_folder( api.upload_folder(
repo_id=hps.repo_id, repo_id=hps.repo_id,
folder_path=model_dir, folder_path=model_dir,
path_in_repo=f"Data/{config.model_name}/models",
delete_patterns="*.pth",
) )
api.upload_folder( api.upload_folder(
repo_id=hps.repo_id, repo_id=hps.repo_id,
folder_path=config.out_dir, folder_path=config.out_dir,
path_in_repo=f"model_assets/{config.model_name}",
) )
if pbar is not None: if pbar is not None:
@@ -937,11 +940,14 @@ def train_and_evaluate(
if hps.repo_id is not None: if hps.repo_id is not None:
api.upload_folder( api.upload_folder(
repo_id=hps.repo_id, repo_id=hps.repo_id,
folder_path=model_dir, folder_path=hps.model_dir,
path_in_repo=f"Data/{config.model_name}/models",
delete_patterns="*.pth",
) )
api.upload_folder( api.upload_folder(
repo_id=hps.repo_id, repo_id=hps.repo_id,
folder_path=config.out_dir, folder_path=config.out_dir,
path_in_repo=f"model_assets/{config.model_name}",
) )
global_step += 1 global_step += 1