mirror of
https://github.com/wassname/Castor.git
synced 2026-08-20 12:00:37 +08:00
* Add Regularization Modules for LSTM * Update Reuters Trainer and Evalueator for regularization * Remove unnecessary comments * Comply with PEP8 * Comply import order with PEP8 * Fix typos in README.md * Comply with PEP8 * Add BSD 3-Clause Licence * Remove deprecated call to Variable for PyTorch 0.4 * Update dataset selection in main * Remove block comments
54 lines
2.1 KiB
Python
54 lines
2.1 KiB
Python
"""
|
|
BSD 3-Clause License
|
|
|
|
Copyright (c) 2017,
|
|
All rights reserved.
|
|
|
|
Redistribution and use in source and binary forms, with or without
|
|
modification, are permitted provided that the following conditions are met:
|
|
|
|
* Redistributions of source code must retain the above copyright notice, this
|
|
list of conditions and the following disclaimer.
|
|
|
|
* Redistributions in binary form must reproduce the above copyright notice,
|
|
this list of conditions and the following disclaimer in the documentation
|
|
and/or other materials provided with the distribution.
|
|
|
|
* Neither the name of the copyright holder nor the names of its
|
|
contributors may be used to endorse or promote products derived from
|
|
this software without specific prior written permission.
|
|
|
|
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
|
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
|
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
|
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
|
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
|
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
|
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
|
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
|
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
|
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
|
"""
|
|
import numpy as np
|
|
|
|
import torch
|
|
|
|
def embedded_dropout(embed, words, dropout=0.1, scale=None):
|
|
if dropout:
|
|
mask = embed.weight.data.new().resize_((embed.weight.size(0), 1)).bernoulli_(1 - dropout).expand_as(embed.weight) / (1 - dropout)
|
|
masked_embed_weight = mask * embed.weight
|
|
else:
|
|
masked_embed_weight = embed.weight
|
|
if scale:
|
|
masked_embed_weight = scale.expand_as(masked_embed_weight) * masked_embed_weight
|
|
|
|
padding_idx = embed.padding_idx
|
|
if padding_idx is None:
|
|
padding_idx = -1
|
|
|
|
X = torch.nn.functional.embedding(words, masked_embed_weight,
|
|
padding_idx, embed.max_norm, embed.norm_type,
|
|
embed.scale_grad_by_freq, embed.sparse
|
|
)
|
|
return X
|