hedgefund

zig_autograd_01

Nov 27th, 2024
263
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C 4.79 KB | Software | 0 0
  1. const std = @import("std");
  2. const testing = std.testing;
  3. const print = std.debug.print;
  4. const time = std.time.milliTimestamp;
  5. const Allocator = std.mem.Allocator;
  6. const ArenaAlloc = std.heap.ArenaAllocator;
  7. const pageAlloc = std.heap.page_allocator;
  8.  
  9. const OpType = enum(u2) {
  10.     Sum,
  11.     Prod,
  12.     Softplus,
  13. };
  14.  
  15. const NaiveVar = struct {
  16.     val: f64,
  17.     grad: f64,
  18.  
  19.     pub fn init(val: f64) NaiveVar {
  20.         return .{ .val = val, .grad = 0.0 };
  21.     }
  22. };
  23.  
  24. const Operation = struct {
  25.     op_type: OpType,
  26.     inputs: [2]*NaiveVar,
  27.     output: NaiveVar,
  28. };
  29.  
  30. const NaiveTape = struct {
  31.     ops: []Operation,
  32.     allocator: Allocator,
  33.  
  34.     pub fn init(allocator: Allocator) NaiveTape {
  35.         return .{
  36.             .ops = &[_]Operation{},
  37.             .allocator = allocator,
  38.         };
  39.     }
  40.  
  41.     pub fn deinit(self: *NaiveTape) void {
  42.         self.allocator.free(self.ops);
  43.     }
  44.  
  45.     pub fn sum(self: *NaiveTape, input1: *NaiveVar, input2: *NaiveVar) !*NaiveVar {
  46.         const sum_val = input1.val + input2.val;
  47.  
  48.         const sum_var = NaiveVar.init(sum_val);
  49.  
  50.         const op = Operation{
  51.             .op_type = .Sum,
  52.             .inputs = [_]*NaiveVar{ input1, input2 },
  53.             .output = sum_var,
  54.         };
  55.  
  56.         self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
  57.         self.ops[self.ops.len - 1] = op;
  58.  
  59.         return &self.ops[self.ops.len - 1].output;
  60.     }
  61.  
  62.     pub fn prod(self: *NaiveTape, var1: *NaiveVar, var2: *NaiveVar) !*NaiveVar {
  63.         const prod_val = var1.val * var2.val;
  64.         const prod_var = NaiveVar.init(prod_val);
  65.  
  66.         const op = Operation{
  67.             .op_type = .Prod,
  68.             .inputs = [_]*NaiveVar{ var1, var2 },
  69.             .output = prod_var,
  70.         };
  71.         self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
  72.         self.ops[self.ops.len - 1] = op;
  73.  
  74.         return &self.ops[self.ops.len - 1].output;
  75.     }
  76.  
  77.     pub fn softplus(self: *NaiveTape, nvar: *NaiveVar) !*NaiveVar {
  78.         const softplus_val = std.math.log1p(std.math.exp(nvar.val));
  79.         const softplus_var = NaiveVar.init(softplus_val);
  80.  
  81.         const op = Operation{
  82.             .op_type = .Softplus,
  83.             .inputs = [_]*NaiveVar{nvar, undefined}, // Only one input needed, second is unused
  84.             .output = softplus_var,
  85.         };
  86.         self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
  87.         self.ops[self.ops.len - 1] = op;
  88.         return &self.ops[self.ops.len - 1].output;
  89.     }
  90.  
  91.     pub fn backward(self: *NaiveTape, nvar: *NaiveVar) void {
  92.         nvar.grad += 1.0;
  93.         var i = self.ops.len;
  94.         while (i > 0) {
  95.             i -= 1;
  96.             const op = self.ops[i];
  97.             switch (op.op_type) {
  98.                 .Sum => {
  99.                     const output_grad = op.output.grad;
  100.                     op.inputs[0].grad += output_grad;
  101.                     op.inputs[1].grad += output_grad;
  102.                 },
  103.                 .Prod => {
  104.                     const output_grad = op.output.grad;
  105.                     const input1_val = op.inputs[0].val;
  106.                     const input2_val = op.inputs[1].val;
  107.                     op.inputs[0].grad += input2_val * output_grad;
  108.                     op.inputs[1].grad += input1_val * output_grad;
  109.                 },
  110.                 .Softplus => {
  111.                     const output_grad = op.output.grad;
  112.                     const input_val = op.inputs[0].val;
  113.                     op.inputs[0].grad += output_grad / (1.0 + std.math.exp(-input_val));
  114.                 },
  115.             }
  116.         }
  117.     }
  118. };
  119.  
  120. pub fn main() !void {
  121.     var arena = ArenaAlloc.init(pageAlloc);
  122.     defer arena.deinit();
  123.     const allocator = arena.allocator();
  124.  
  125.     const iterations: usize = 1000000;
  126.  
  127.     const start_time = time();
  128.     var i: usize = 0;
  129.     while (i < iterations) : (i += 1) {
  130.         var var1 = NaiveVar.init(1.0);
  131.         var var2 = NaiveVar.init(2.0);
  132.  
  133.         var tape = NaiveTape.init(allocator);
  134.         defer tape.deinit();
  135.  
  136.         const sum_var = try tape.sum(&var1, &var2);
  137.         const prod_var = try tape.prod(sum_var, sum_var);
  138.         const softplus_var = try tape.softplus(prod_var);
  139.  
  140.         tape.backward(softplus_var);
  141.  
  142.         if (i == iterations - 1) {
  143.             print("sum_var val: {d:}\n", .{sum_var.val});
  144.             print("prod_var val: {d:}\n", .{prod_var.val});
  145.             print("softplus_var val: {d:}\n", .{softplus_var.val});
  146.             print("sum_var grad: {d:}\n", .{sum_var.grad});
  147.             print("var1 grad: {d:}\n", .{var1.grad});
  148.             print("var2 grad: {d:}\n", .{var2.grad});
  149.         }
  150.     }
  151.     const end_time = time();
  152.     const elapsed_time = @as(f64, @floatFromInt(end_time - start_time));
  153.     print("\nElapsed time: {d:.3} ms\n", .{elapsed_time});
  154. }
  155.  
Advertisement
Add Comment
Please, Sign In to add comment