graykode / graykode/nlp-tutorial

about seq2seq(attention)-Torch multiple sample training question

Ouverte
#54 0 commentaires 0 réactions 0 personnes assignées Voir sur GitHub
Langage dominant
Jupyter Notebook
Étoiles
14.9k
Forks
3.9k
Métriques de merge des PR
Aucune PR mergée en 30 j

Description

hello, first thank your code, but i want to know if batch_size is more than 1, i should how to modify the code, thank you
```
def get_att_weight(self, output, enc_output): # get attention weight one 'output' with 'enc_output'
'''
output: [1, batch_size, num_directions(=1) * n_hidden]
enc_output: [n_step+1, batch_size, num_directions(=1) * n_hidden]
'''
length = len(enc_output)
attn_scores = torch.zeros(length) # attn_scores : [batch_size, n_step+1]
for i in range(length):
attn_scores[i] = self.get_att_score(output, enc_output[i])

# Normalize scores to weights in range 0 to 1
# return [batch_size, 1, n_step+1]
return F.softmax(attn_scores).view(batch_size, 1, -1)

def get_att_score(self, output, enc_output):
'''
output: [batch_size, num_directions(=1) * n_hidden]
enc_output: [batch_size, num_directions(=1) * n_hidden]
'''
score = self.attn(enc_output) # score : [1, n_hidden]
return torch.dot(output.view(-1), score.view(-1)) # inner product make scalar value, get a real number
```

Guide de contribution

Ouvrir le guide de contribution

Piste de recherche

Commencez par les méthodes fournies get_att_weight et get_att_score et examinez l’implémentation seq2seq de l’attention qui les entoure. Suivez les formes des tenseurs pour un batch_size supérieur à un, puis vérifiez que l’entraînement de l’attention fonctionne pour plusieurs échantillons sans erreurs de forme scalaire. La tâche est terminée lorsque l’exemple prend en charge batch_size > 1 avec des poids d’attention correctement dimensionnés.

Rédigé par le modèle d'indexation à partir du texte de l'issue.

Évaluation

Stack technique
python, pytorch
Domaine
machine-learning
Type d'issue
Fonctionnalité
Difficulté
4/5
Temps estimé
3-5 jours
Activité
À l'abandon
Clarté
Plutôt claire
Accessibilité débutants
35/100

Recevez les nouvelles issues par e-mail

Un résumé court des issues GitHub adaptées aux débutants.