博客
关于我
神经网络中的矩阵求导及反向传播推导
阅读量:321 次
发布时间:2019-03-03

本文共 1475 字,大约阅读时间需要 4 分钟。

两层全连接神经网络的实现, 包括网络的实现、梯度的反向传播计算和权重更新过程:

 

# -*- coding: utf-8 -*-import numpy as np# N is batch size; D_in is input dimension;# H is hidden dimension; D_out is output dimension.N, D_in, H, D_out = 64, 1000, 100, 10# Create random input and output datax = np.random.randn(N, D_in)y = np.random.randn(N, D_out)# Randomly initialize weightsw1 = np.random.randn(D_in, H)w2 = np.random.randn(H, D_out)learning_rate = 1e-6for t in range(500):    # Forward pass: compute predicted y    h = x.dot(w1)    h_relu = np.maximum(h, 0)    y_pred = h_relu.dot(w2)    # Compute and print loss    loss = np.square(y_pred - y).sum()    print(t, loss)    # Backprop to compute gradients of w1 and w2 with respect to loss    grad_y_pred = 2.0 * (y_pred - y)    grad_w2 = h_relu.T.dot(grad_y_pred)    grad_h_relu = grad_y_pred.dot(w2.T)    grad_h = grad_h_relu.copy()    grad_h[h < 0] = 0    grad_w1 = x.T.dot(grad_h)    # Update weights    w1 -= learning_rate * grad_w1    w2 -= learning_rate * grad_w2

这里解决了我一个错误的认知:以为最速下降法跟各个变量计算的导数无关, 而其实就是每个变量各自按自己的导数下降就可以实现函数最陡的坡进行下降;在图形上可以理解多个向量合并成一个方向;

反向传播 过程,核心代码如下

h = x.dot(w1)h_relu = np.maximum(h, 0)y_pred = h_relu.dot(w2)loss = np.square(y_pred - y).sum()grad_y_pred = 2.0 * (y_pred - y)    # 64 x 10grad_w2 = h_relu.T.dot(grad_y_pred) # 100 x 10grad_h_relu = grad_y_pred.dot(w2.T) # 64 x 100grad_h = grad_h_relu.copy()         # 64 x 100grad_h[h < 0] = 0                   # 64 x 100grad_w1 = x.T.dot(grad_h)           # 1000 x 100

 

 

 

 

 

问题:如何实现relu求导呢?

转载地址:http://usgm.baihongyu.com/

你可能感兴趣的文章
mysql 的存储引擎介绍
查看>>
MySQL 的存储引擎有哪些?为什么常用InnoDB?
查看>>
Mysql 知识回顾总结-索引
查看>>
Mysql 笔记
查看>>
MySQL 精选 60 道面试题(含答案)
查看>>
mysql 索引
查看>>
MySQL 索引失效的 15 种场景!
查看>>
MySQL 索引深入解析及优化策略
查看>>
MySQL 索引的面试题总结
查看>>
mysql 索引类型以及创建
查看>>
MySQL 索引连环问题,你能答对几个?
查看>>
Mysql 索引问题集锦
查看>>
Mysql 纵表转换为横表
查看>>
mysql 编译安装 window篇
查看>>
mysql 网络目录_联机目录数据库
查看>>
MySQL 聚簇索引&&二级索引&&辅助索引
查看>>
Mysql 脏页 脏读 脏数据
查看>>
mysql 自增id和UUID做主键性能分析,及最优方案
查看>>
Mysql 自定义函数
查看>>
mysql 行转列 列转行
查看>>