Update train_ms.py
This commit is contained in:
@@ -7,8 +7,8 @@ import torch
|
|||||||
from torch import nn, optim
|
from torch import nn, optim
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
#from torch.utils.tensorboard import SummaryWriter
|
||||||
#from tensorboardX import SummaryWriter
|
from wandb.tensorboard import WandbSummaryWriter as SummaryWriter
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
@@ -190,7 +190,7 @@ def run(rank, n_gpus, hps):
|
|||||||
scheduler_g.step()
|
scheduler_g.step()
|
||||||
scheduler_d.step()
|
scheduler_d.step()
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
cheduler_dur_disc.step()
|
scheduler_dur_disc.step()
|
||||||
|
|
||||||
|
|
||||||
def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loaders, logger, writers):
|
def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loaders, logger, writers):
|
||||||
|
|||||||
Reference in New Issue
Block a user