Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
429c676f95 |
@@ -110,12 +110,13 @@ def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int]
|
||||
|
||||
|
||||
def BartLayerNorm(normalized_shape: torch.Size, eps: float = 1e-5, elementwise_affine: bool = True):
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
if torch.cuda.is_available():
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
|
||||
return FusedLayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
except ImportError:
|
||||
pass
|
||||
return FusedLayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
except ImportError:
|
||||
pass
|
||||
return torch.nn.LayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
|
||||
|
||||
|
||||
@@ -265,12 +265,14 @@ FSMT_INPUTS_DOCSTRING = r"""
|
||||
|
||||
|
||||
have_fused_layer_norm = False
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
if torch.cuda.is_available():
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
|
||||
have_fused_layer_norm = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
have_fused_layer_norm = True
|
||||
except ImportError:
|
||||
pass
|
||||
LayerNorm = FusedLayerNorm if have_fused_layer_norm else torch.nn.LayerNorm
|
||||
|
||||
|
||||
|
||||
@@ -511,12 +511,13 @@ class ProphetNetDecoderLMOutput(ModelOutput):
|
||||
|
||||
|
||||
def ProphetNetLayerNorm(normalized_shape, eps=1e-5, elementwise_affine=True):
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
if torch.cuda.is_available():
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
|
||||
return FusedLayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
except ImportError:
|
||||
pass
|
||||
return FusedLayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
except ImportError:
|
||||
pass
|
||||
return torch.nn.LayerNorm(normalized_shape, eps, elementwise_affine)
|
||||
|
||||
|
||||
|
||||
@@ -228,8 +228,9 @@ class Trainer:
|
||||
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),
|
||||
):
|
||||
if args is None:
|
||||
logger.info("No `TrainingArguments` passed, using the current path as `output_dir`.")
|
||||
args = TrainingArguments("tmp_trainer")
|
||||
output_dir = "tmp_trainer"
|
||||
logger.info(f"No `TrainingArguments` passed, using `output_dir={output_dir}`.")
|
||||
args = TrainingArguments(output_dir=output_dir)
|
||||
self.args = args
|
||||
# Seed must be set before instantiating the model when using model
|
||||
set_seed(self.args.seed)
|
||||
|
||||
Reference in New Issue
Block a user