Fix hf push
This commit is contained in:
27
train_ms.py
27
train_ms.py
@@ -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:
|
||||||
|
|||||||
@@ -574,15 +574,18 @@ def run():
|
|||||||
),
|
),
|
||||||
for_infer=True,
|
for_infer=True,
|
||||||
)
|
)
|
||||||
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=model_dir,
|
||||||
)
|
path_in_repo=f"Data/{config.model_name}/models",
|
||||||
api.upload_folder(
|
delete_patterns="*.pth",
|
||||||
repo_id=hps.repo_id,
|
)
|
||||||
folder_path=config.out_dir,
|
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()
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user