fixing Win failed import (#1163)

* version

* try fix distrib

* update try import
This commit is contained in:
Jirka Borovec
2020-03-16 20:50:36 -04:00
committed by GitHub
parent 49d000c0c9
commit e461ec0037
8 changed files with 18 additions and 19 deletions
@@ -1,13 +1,12 @@
import logging as log
from abc import ABC
try:
from apex import amp
APEX_AVAILABLE = True
except ImportError:
APEX_AVAILABLE = False
import logging as log
else:
APEX_AVAILABLE = True
class TrainerAMPMixin(ABC):
+2 -2
View File
@@ -1,7 +1,7 @@
from abc import ABC, abstractmethod
from typing import Union, List, Tuple, Callable
import torch.distributed as dist
import torch.distributed as torch_distrib
from torch.utils.data import SequentialSampler, DataLoader
from torch.utils.data.distributed import DistributedSampler
@@ -224,7 +224,7 @@ class TrainerDataLoadingMixin(ABC):
# get the function we'll use to get data
if self.use_ddp or self.use_ddp2:
# all processes wait until data download has happened
dist.barrier()
torch_distrib.barrier()
# data download/load on TPU
elif self.use_tpu and XLA_AVAILABLE:
+2 -2
View File
@@ -8,7 +8,7 @@ from typing import Union, Optional, List, Dict, Tuple, Iterable
import torch
from torch import optim
import torch.distributed as dist
import torch.distributed as torch_distrib
import torch.multiprocessing as mp
from torch.optim.optimizer import Optimizer
from torch.utils.data import DataLoader
@@ -748,7 +748,7 @@ class Trainer(
self.logger.save()
if self.use_ddp or self.use_ddp2:
dist.barrier()
torch_distrib.barrier()
# wait for all models to restore weights
if self.on_tpu and XLA_AVAILABLE:
+2 -2
View File
@@ -100,7 +100,7 @@ from subprocess import call
from typing import Union
import torch
import torch.distributed as dist
import torch.distributed as torch_distrib
from pytorch_lightning.core.lightning import LightningModule
from pytorch_lightning.loggers import LightningLoggerBase
@@ -177,7 +177,7 @@ class TrainerIOMixin(ABC):
# wait for all models to restore weights
if self.use_ddp or self.use_ddp2:
# wait for all processes to catch up
dist.barrier()
torch_distrib.barrier()
# wait for all models to restore weights
if self.on_tpu and XLA_AVAILABLE: