MCPcopy Create free account

hub / github.com/Doragd/Chinese-Chatbot-PyTorch-Implementation / functions

Functions31 in github.com/Doragd/Chinese-Chatbot-PyTorch-Implementation

↓ 3 callersFunctionget_dataloader
(opt)
dataload.py:99
↓ 2 callersMethod__init__
(self, attn_method, hidden_size)
model.py:60
↓ 2 callersFunctiongenerate
(input_seq, searcher, sos, eos, opt)
train_eval.py:194
↓ 2 callersFunctionmaskNLLLoss
inp: shape [batch_size,voc_length] target: shape [batch_size] 经过view ==> [batch_size, 1] 这样就和inp维数相同,可以用gather target作为索引,在dim=1上索
train_eval.py:16
↓ 2 callersFunctionpreprocess
()
datapreprocess.py:20
↓ 2 callersFunctionzeroPadding
l是多个长度不同的句子(list),使用zip_longest padding成定长,长度为最长句子的长度 在zeroPadding函数中隐式转置 [batch_size, max_seq_len] ==> [max_seq_len, batch_size]
dataload.py:7
↓ 1 callersFunctionbinaryMatrix
生成mask矩阵, 0表示padding,1表示未padding shape同l,即[max_seq_len, batch_size]
dataload.py:15
↓ 1 callersMethodconcat_score
hidden: h_t, shape: [max_seq_len, batch_size, hidden_size] expand(max_seq_len, -1,-1) ==> [max_seq_len, batch_size
model.py:98
↓ 1 callersFunctioncreate_collate_fn
说明dataloader如何包装一个batch,传入的参数为</PAD>的索引padding,</EOS>字符索引eos collate_fn传入的参数是由一个batch的__getitem__方法的返回值组成的corpus_item corpus_item:
dataload.py:30
↓ 1 callersMethoddot_score
encoder_outputs: encoder(双向GRU)的所有时刻的最后一层的hidden输出 shape: [max_seq_len, batch_size, hidden_size] 数学符号
model.py:72
↓ 1 callersMethodgeneral_score
(self, hidden, encoder_outputs)
model.py:92
↓ 1 callersFunctiontrain_by_batch
(sos, opt, data, encoder_optimizer, decoder_optimizer, encoder, decoder)
train_eval.py:33
↓ 1 callersFunctionupdate
(word_nums)
datapreprocess.py:37
Method__getitem__
(self, index)
dataload.py:90
Method__init__
(self, opt)
dataload.py:81
Method__init__
voc_length: 字典长度,即输入的单词的one-hot编码长度
model.py:10
Method__init__
(self, opt, voc_length)
model.py:124
Method__init__
(self, encoder, decoder)
utils/greedysearch.py:9
Method__len__
(self)
dataload.py:95
Functionchat
(**kwargs)
main.py:11
Functioncollate_fn
(corpus_item)
dataload.py:57
Functioneval
(**kwargs)
train_eval.py:205
Methodforward
input_seq: shape: [max_seq_len, batch_size] input_lengths: 一批次中每个句子对应的句子长度列表 shape:[batch_
model.py:23
Methodforward
(self, hidden, encoder_outputs)
model.py:111
Methodforward
input_step: decoder是逐字生成的,即每个timestep产生一个字, decoder接收的输入: input_step='/SOS'的索引 和 encoder的最后时刻的最后一层hidden输出
model.py:139
Methodforward
(self, sos, eos, input_seq, input_length, max_length, device)
utils/greedysearch.py:14
Functionfun
(word)
datapreprocess.py:38
Functionmatch
(input_question)
QA_data/QA_test.py:16
Functionoutput_answer
(input_sentence, searcher, sos, eos, unknown, opt, word2ix, ix2word)
train_eval.py:286
Functiontest
(opt)
train_eval.py:253
Functiontrain
(**kwargs)
train_eval.py:130