You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

二、短密码处理逻辑

  1. Padding与Packed Sequence机制

    • 所有密码统一padding到最大长度30,用<placeholder>作为填充字符(见第三部分)
    • 使用pack_padded_sequence将padding部分从GRU计算中排除,GRU仅处理每个密码的真实有效字符,剩余padding单元完全不参与计算
    • 预处理时需记录每个密码对的真实长度(取原始/当前密码的较长值),forward时传入该参数即可
  2. 批量处理注意事项

    • 无需手动对样本按长度排序,enforce_sorted=False可自动处理乱序的长度输入

三、特殊字符集成方法

  1. 扩展字符词汇表

    • 在基础字符表中加入<placeholder>、<shift>、<caps>三个特殊字符,分配唯一索引(例如放在字符表首尾),示例结构:
      vocab = ['<placeholder>', '<shift>', '<caps>', 'a', 'b', ..., 'Z', '0', ..., '!']
      
    • 确保模型Embedding层的vocab_size包含这些特殊字符
  2. 输入编码规则

    • <placeholder>:用于填充短密码至30长度,例如长度为8的密码,后续补22个<placeholder>
    • <shift>/<caps>:对应论文中键盘操作建模需求,预处理时将编辑行为映射为特殊字符插入:
      • 如原始密码abc、当前密码Abc,需在原始密码对应位置前插入<caps>标记;若为aBc则插入<shift>标记
      • 严格遵循论文3.2节的输入编码规则,将特殊字符融入原始/当前密码的序列中

四、多分类无关类处理

  1. 训练阶段:损失掩码

    • 对每个样本预计算不可能的类别(例如长度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
      
  2. 推理阶段:过滤无效类

    • 根据输入密码长度生成无效类列表,将对应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
      

内容的提问来源于stack exchange,提问作者srswat

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 12:53:10