码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • NNDL 作业8:RNN - 简单循环网络


    简单循环网络 ( Simple Recurrent Network , SRN) 只有一个隐藏层的神经网络 .

    目录

    1. 使用Numpy实现SRN

    2. 在1的基础上,增加激活函数tanh

    3. 分别使用nn.RNNCell、nn.RNN实现SRN

    4. 分析“二进制加法” 源代码(选做)

    5. 实现“Character-Level Language Models”源代码(必做)

    6. 分析“序列到序列”源代码(选做)

    7. “编码器-解码器”的简单实现(必做)


    1. 使用Numpy实现SRN

    1. import numpy as np
    2. inputs = np.array([[1., 1.],
    3. [1., 1.],
    4. [2., 2.]]) # 初始化输入序列
    5. print('inputs is ', inputs)
    6. state_t = np.zeros(2, ) # 初始化存储器
    7. print('state_t is ', state_t)
    8. w1, w2, w3, w4, w5, w6, w7, w8 = 1., 1., 1., 1., 1., 1., 1., 1.
    9. U1, U2, U3, U4 = 1., 1., 1., 1.
    10. print('--------------------------------------')
    11. for input_t in inputs:
    12. print('inputs is ', input_t)
    13. print('state_t is ', state_t)
    14. in_h1 = np.dot([w1, w3], input_t) + np.dot([U2, U4], state_t)
    15. in_h2 = np.dot([w2, w4], input_t) + np.dot([U1, U3], state_t)
    16. state_t = in_h1, in_h2
    17. output_y1 = np.dot([w5, w7], [in_h1, in_h2])
    18. output_y2 = np.dot([w6, w8], [in_h1, in_h2])
    19. print('output_y is ', output_y1, output_y2)
    20. print('---------------')

    2. 在1的基础上,增加激活函数tanh

    1. import numpy as np
    2. inputs = np.array([[1., 1.],
    3. [1., 1.],
    4. [2., 2.]]) # 初始化输入序列
    5. print('inputs is ', inputs)
    6. state_t = np.zeros(2, ) # 初始化存储器
    7. print('state_t is ', state_t)
    8. w1, w2, w3, w4, w5, w6, w7, w8 = 1., 1., 1., 1., 1., 1., 1., 1.
    9. U1, U2, U3, U4 = 1., 1., 1., 1.
    10. print('--------------------------------------')
    11. for input_t in inputs:
    12. print('inputs is ', input_t)
    13. print('state_t is ', state_t)
    14. in_h1 = np.tanh(np.dot([w1, w3], input_t) + np.dot([U2, U4], state_t))
    15. in_h2 = np.tanh(np.dot([w2, w4], input_t) + np.dot([U1, U3], state_t))
    16. state_t = in_h1, in_h2
    17. output_y1 = np.dot([w5, w7], [in_h1, in_h2])
    18. output_y2 = np.dot([w6, w8], [in_h1, in_h2])
    19. print('output_y is ', output_y1, output_y2)
    20. print('---------------')

    3. 分别使用nn.RNNCell、nn.RNN实现SRN

    1. import torch
    2. batch_size = 1
    3. seq_len = 3 # 序列长度
    4. input_size = 2 # 输入序列维度
    5. hidden_size = 2 # 隐藏层维度
    6. output_size = 2 # 输出层维度
    7. # RNNCell
    8. cell = torch.nn.RNNCell(input_size=input_size, hidden_size=hidden_size)
    9. # 初始化参数 https://zhuanlan.zhihu.com/p/342012463
    10. for name, param in cell.named_parameters():
    11. if name.startswith("weight"):
    12. torch.nn.init.ones_(param)
    13. else:
    14. torch.nn.init.zeros_(param)
    15. # 线性层
    16. liner = torch.nn.Linear(hidden_size, output_size)
    17. liner.weight.data = torch.Tensor([[1, 1], [1, 1]])
    18. liner.bias.data = torch.Tensor([0.0])
    19. seq = torch.Tensor([[[1, 1]],
    20. [[1, 1]],
    21. [[2, 2]]])
    22. hidden = torch.zeros(batch_size, hidden_size)
    23. output = torch.zeros(batch_size, output_size)
    24. for idx, input in enumerate(seq):
    25. print('=' * 20, idx, '=' * 20)
    26. print('Input :', input)
    27. print('hidden :', hidden)
    28. hidden = cell(input, hidden)
    29. output = liner(hidden)
    30. print('output :', output)

    1. import torch
    2. batch_size = 1
    3. seq_len = 3
    4. input_size = 2
    5. hidden_size = 2
    6. num_layers = 1
    7. output_size = 2
    8. cell = torch.nn.RNN(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers)
    9. for name, param in cell.named_parameters(): # 初始化参数
    10. if name.startswith("weight"):
    11. torch.nn.init.ones_(param)
    12. else:
    13. torch.nn.init.zeros_(param)
    14. # 线性层
    15. liner = torch.nn.Linear(hidden_size, output_size)
    16. liner.weight.data = torch.Tensor([[1, 1], [1, 1]])
    17. liner.bias.data = torch.Tensor([0.0])
    18. inputs = torch.Tensor([[[1, 1]],
    19. [[1, 1]],
    20. [[2, 2]]])
    21. hidden = torch.zeros(num_layers, batch_size, hidden_size)
    22. out, hidden = cell(inputs, hidden)
    23. print('Input :', inputs[0])
    24. print('hidden:', 0, 0)
    25. print('Output:', liner(out[0]))
    26. print('--------------------------------------')
    27. print('Input :', inputs[1])
    28. print('hidden:', out[0])
    29. print('Output:', liner(out[1]))
    30. print('--------------------------------------')
    31. print('Input :', inputs[2])
    32. print('hidden:', out[1])
    33. print('Output:', liner(out[2]))

    4. 分析“二进制加法” 源代码(选做)

    Anyone Can Learn To Code an LSTM-RNN in Python (Part 1: RNN) - i am trask

    5. 实现“Character-Level Language Models”源代码(必做)

    翻译Character-Level Language Models 相关内容

    The Unreasonable Effectiveness of Recurrent Neural Networks

    编码实现该模型 

    6. 分析“序列到序列”源代码(选做)

     

    7. “编码器-解码器”的简单实现(必做)

     

    seq2seq的PyTorch实现_哔哩哔哩_bilibili

    Seq2Seq的PyTorch实现 - mathor

    REF:

    Hung-yi Lee (ntu.edu.tw)

    《PyTorch深度学习实践》完结合集_哔哩哔哩_bilibili

    完全图解RNN、RNN变体、Seq2Seq、Attention机制 - 知乎 (zhihu.com)
  • 相关阅读:
    onlyoffice的介绍搭建、集成过程。Windows、Linux
    JDK19新特性使用详解
    Python 有哪些好的学习资料或者博客?
    c++实现观察者模式
    【自学开发之旅】Flask-标准化返回-连接数据库-分表-orm-migrate-增删改查(三)
    Qt 打印调试信息-怎样获取QTableWidget的行数和列数-读取QTableWidget表格中的数据
    MySQL之索引初识篇:索引机制、索引分类、索引使用与管理综述
    红旗系统Asianux 8.1常用命令(配置jdk、mysql、redis、RabbitMQ等等)
    2023年电工杯 | 2023年电工杯数学建模数学建模竞赛思路(A题、B题)
    「精致店主理人」:青年并肩同行,共赴创业之路
  • 原文地址:https://blog.csdn.net/qq_38975453/article/details/127561213
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号