ND4J自动微分

一、前言

    ND4J从beta2开始就开始支持自动微分,不过直到beta4版本为止,自动微分还只支持CPU,GPU版本将在后续版本中实现。

    本篇博客中,我们将用ND4J来构建一个函数,利用ND4J SameDiff构建函数求函数值和求函数每个变量的偏微分值。

二、构建函数

    构建函数和分别手动求偏导数

    

    给定一个点(2,3)手动求函数值和偏导,计算如下:

    f=2+3*4+3=17,f对x的偏导:1+2*2*3=13,f对y的偏导:4+1=5

三、通过ND4J自动微分来求

    完整代码

1package org.nd4j.samediff; 2 3import org.nd4j.autodiff.samediff.SDVariable; 4import org.nd4j.autodiff.samediff.SameDiff; 5import org.nd4j.linalg.factory.Nd4j; 6 7/** 8 * 9 * x+y*x2+y 10 * 11 */ 12public class Function { 13 14 public static void main(String[] args) { 15 //构建SameDiff实例 16 SameDiff sd=SameDiff.create(); 17 //创建变量x、y 18 SDVariable x= sd.var("x"); 19 SDVariable y=sd.var("y"); 20 21 //定义函数 22 SDVariable f=x.add(y.mul(sd.math().pow(x, 2))); 23 f.add("addY",y); 24 25 //给变量x、y绑定具体值 26 x.setArray(Nd4j.create(new double[]{2})); 27 y.setArray(Nd4j.create(new double[]{3})); 28 //前向计算函数的值 29 System.out.println(sd.exec(null, "addY").get("addY")); 30 //后向计算求梯度 31 sd.execBackwards(null); 32 //打印x在(2,3)处的导数 33 System.out.println(sd.getGradForVariable("x").getArr()); 34 //x.getGradient().getArr()和sd.getGradForVariable("x").getArr()等效 35 System.out.println(x.getGradient().getArr()); 36 //打印y在(2,3)处的导数 37 System.out.println(sd.getGradForVariable("y").getArr()); 38 } 39}

    四、运行结果

1o.n.l.f.Nd4jBackend - Loaded [CpuBackend] backend 2o.n.n.NativeOpsHolder - Number of threads used for NativeOps: 4 3o.n.n.Nd4jBlas - Number of threads used for BLAS: 4 4o.n.l.a.o.e.DefaultOpExecutioner - Backend used: [CPU]; OS: [Windows 10] 5o.n.l.a.o.e.DefaultOpExecutioner - Cores: [8]; Memory: [3.2GB]; 6o.n.l.a.o.e.DefaultOpExecutioner - Blas vendor: [MKL] 717.0000 8o.n.a.s.SameDiff - Inferring output "addY" as loss variable as none were previously set. Use SameDiff.setLossVariables() to override 913.0000 1013.0000 115.0000

    结果为17、13、5和手动求出的结果完全一致。

    自动微分屏蔽了deeplearning在求微分过程中的很多细节,特别是矩阵求导、矩阵范数求导等等,是非常麻烦的,用自动微分,可以轻松实现各式各样的网络结构。

快乐源于分享。

此博客乃作者原创, 转载请注明出处

点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

皕杰报表之UUID

​在我们用皕杰报表工具设计填报报表时,如何在新增行里自动增加id呢?能新增整数排序id吗?目前可以在新增行里自动增加id,但只能用uuid函数增加UUID编码,不能新增整数排序id。uuid函数说明:获取一个UUID,可以在填报表中用来创建数据ID语法:uuid()或uuid(sep)参数说明:sep布尔值,生成的uuid中是否包含分隔符'',缺省为

手写Java HashMap源码

HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程22

Lua基础(对象)

:和.区别.   stu{id100,name"Tom",age21}成员变量   function stu.toString()成员函数    return stu.id .. stu.name .. stu.age   endprint(stu

JS 对象数组Array 根据对象object key的值排序sort,很风骚哦

有个js对象数组varary\{id:1,name:"b"},{id:2,name:"b"}\需求是根据name或者id的值来排序,这里有个风骚的函数函数定义:function keysrt(key,desc) {  return function(a,b){    return desc ? ~~(ak