12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091(* Claude Code
*
* Copyright (C) 2026 Yoann Padioleau
*
* This library is free software; you can redistribute it and/or
* modify it under the terms of the GNU Library General Public License
* (LGPL) as published by the Free Software Foundation; either version
* 2 of the License, or (at your option) any later version.
*)(* See Grad.mli *)(* a node: its value, the slope of the final answer with respect to
* it, what it was made from, and how to send a slope back to those --
* which is the only thing an operation has to know about itself *)typet={mutablev:float;mutabled:float;(* the slope, filled in by [backward] *)from:tlist;send_back:t->unit;(* takes the node, pushes its d onto [from] *)}letnothing(_:t):unit=()letvalue(v:float):t={v;d=0.;from=[];send_back=nothing}letof_(n:t):float=n.vletslope(n:t):float=n.dletmake(v:float)(from:tlist)(send_back:t->unit):t={v;d=0.;from;send_back}(* every operation is its value and where its slope goes. The sums are
* the chain rule: a node's slope is added to each input's, multiplied
* by that input's local derivative *)let(+:)(a:t)(b:t):t=make(a.v+.b.v)[a;b](funn->a.d<-a.d+.n.d;b.d<-b.d+.n.d)let(-:)(a:t)(b:t):t=make(a.v-.b.v)[a;b](funn->a.d<-a.d+.n.d;b.d<-b.d-.n.d)let(*:)(a:t)(b:t):t=make(a.v*.b.v)[a;b](funn->(* each input's slope is the other input: d(ab)/da = b *)a.d<-a.d+.(n.d*.b.v);b.d<-b.d+.(n.d*.a.v))let(/:)(a:t)(b:t):t=make(a.v/.b.v)[a;b](funn->a.d<-a.d+.(n.d/.b.v);b.d<-b.d-.(n.d*.a.v/.(b.v*.b.v)))letneg(a:t):t=make(-.a.v)[a](funn->a.d<-a.d-.n.d)letexp_(a:t):t=make(expa.v)[a](funn->a.d<-a.d+.(n.d*.n.v))letlog_(a:t):t=make(loga.v)[a](funn->a.d<-a.d+.(n.d/.a.v))(* the squashes, with the derivatives Net.slope also uses: taken from
* the output, which the node already holds *)lettanh_(a:t):t=make(tanha.v)[a](funn->a.d<-a.d+.(n.d*.(1.-.(n.v*.n.v))))letsigmoid(a:t):t=make(1./.(1.+.exp(-.a.v)))[a](funn->a.d<-a.d+.(n.d*.n.v*.(1.-.n.v)))letrelu(a:t):t=make(ifa.v>0.thena.velse0.)[a](funn->a.d<-a.d+.(ifa.v>0.thenn.delse0.))letsquare(a:t):t=make(a.v*.a.v)[a](funn->a.d<-a.d+.(n.d*.2.*.a.v))letsum(l:tlist):t=List.fold_left(+:)(value0.)l(* the nodes, deepest first: a node may only send its slope back once
* everything it feeds has sent to it, so they are ordered by what
* depends on what *)letorder(root:t):tlist=letseen=ref[]andout=ref[]inletrecgo(n:t)=ifnot(List.memqn!seen)then(seen:=n::!seen;List.itergon.from;out:=n::!out)ingoroot;!outletbackward(root:t):unit=letnodes=orderrootinList.iter(funn->n.d<-0.)nodes;root.d<-1.;List.iter(funn->n.send_backn)nodesletzero(root:t):unit=List.iter(funn->n.d<-0.)(orderroot)letnodes(root:t):int=List.length(orderroot)