Pass2Edit模型PyTorch GRU实现问询:变长密码与特殊字符处理
Pass2Edit PyTorch实现指南
一、核心模型框架实现
直接基于你理解的流程,给出PyTorch代码实现:
import torch import torch.nn as nn import torch.nn.functional as F class Pass2Edit(nn.Module): def __init__(self, vocab_size, embed_dim=256, gru_hidden=256, fc_hidden=512, num_classes=1561): super().__init__() # Embedding层:替代独热编码,将字符索引转为256维向量 self.embedding = nn.Embedding(vocab_size, embed_dim) # 3层单向GRU,batch_first=True适配批量输入格式 self.gru = nn.GRU( input_size=embed_dim * 2, # 原始+当前密码的字符向量拼接 hidden_size=gru_hidden, num_layers=3, batch_first=True ) # 全连接层 self.fc1 = nn.Linear(gru_hidden, fc_hidden) self.fc2 = nn.Linear(fc_hidden, num_classes) self.dropout = nn.Dropout(0.2) # 可选,抑制过拟合 def forward(self, raw_pw, curr_pw, lengths): # raw_pw/curr_pw: [batch_size, max_len],字符索引序列 # lengths: [batch_size],每个样本的真实有效长度(不含padding) # 1. 生成字符嵌入向量 raw_embed = self.embedding(raw_pw) # [batch, max_len, 256] curr_embed = self.embedding(curr_pw) # [batch, max_len, 256] concat_embed = torch.cat([raw_embed, curr_embed], dim=-1) # [batch, max_len, 512] # 2. 处理短密码:跳过padding部分的GRU计算 packed_embed = nn.utils.rnn.pack_padded_sequence( concat_embed, lengths, batch_first=True, enforce_sorted=False ) _, h_n = self.gru(packed_embed) # h_n: [num_layers, batch, 256] # 取最后一层GRU的隐藏状态作为输出 gru_out = h_n[-1] # [batch, 256] # 3. 全连接层输出分类概率 x = self.dropout(F.relu(self.fc1(gru_out))) logits = self.fc2(x) # [batch, 1561] return logits
二、短密码处理逻辑
Padding与Packed Sequence机制
- 所有密码统一padding到最大长度30,用
<placeholder>作为填充字符(见第三部分) - 使用
pack_padded_sequence将padding部分从GRU计算中排除,GRU仅处理每个密码的真实有效字符,剩余padding单元完全不参与计算 - 预处理时需记录每个密码对的真实长度(取原始/当前密码的较长值),forward时传入该参数即可
- 所有密码统一padding到最大长度30,用
批量处理注意事项
- 无需手动对样本按长度排序,
enforce_sorted=False可自动处理乱序的长度输入
- 无需手动对样本按长度排序,
三、特殊字符集成方法
扩展字符词汇表
- 在基础字符表中加入
<placeholder>、<shift>、<caps>三个特殊字符,分配唯一索引(例如放在字符表首尾),示例结构:vocab = ['<placeholder>', '<shift>', '<caps>', 'a', 'b', ..., 'Z', '0', ..., '!'] - 确保模型Embedding层的
vocab_size包含这些特殊字符
- 在基础字符表中加入
输入编码规则
<placeholder>:用于填充短密码至30长度,例如长度为8的密码,后续补22个<placeholder><shift>/<caps>:对应论文中键盘操作建模需求,预处理时将编辑行为映射为特殊字符插入:- 如原始密码
abc、当前密码Abc,需在原始密码对应位置前插入<caps>标记;若为aBc则插入<shift>标记 - 严格遵循论文3.2节的输入编码规则,将特殊字符融入原始/当前密码的序列中
- 如原始密码
四、多分类无关类处理
训练阶段:损失掩码
- 对每个样本预计算不可能的类别(例如长度8的密码,所有位置>8的INS/DEL操作类),生成掩码矩阵(有效类设1,无效类设0)
- 自定义损失函数,仅计算有效类的交叉熵损失:
def custom_loss(logits, targets, mask): log_probs = F.log_softmax(logits, dim=-1) target_probs = log_probs.gather(1, targets.unsqueeze(1)).squeeze() # 仅保留有效类的损失计算 loss = -torch.mean(target_probs * mask.gather(1, targets.unsqueeze(1)).squeeze()) return loss
推理阶段:过滤无效类
- 根据输入密码长度生成无效类列表,将对应logits设为负无穷后再做Softmax,确保无效类概率置0:
def predict(model, raw_pw, curr_pw, length, invalid_classes): logits = model(raw_pw, curr_pw, [length]) logits[0, invalid_classes] = -float('inf') probs = F.softmax(logits, dim=-1) return probs
- 根据输入密码长度生成无效类列表,将对应logits设为负无穷后再做Softmax,确保无效类概率置0:
内容的提问来源于stack exchange,提问作者srswat
相关产品推荐
相关产品推荐

