Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- const std = @import("std");
- const testing = std.testing;
- const print = std.debug.print;
- const time = std.time.milliTimestamp;
- const Allocator = std.mem.Allocator;
- const ArenaAlloc = std.heap.ArenaAllocator;
- const pageAlloc = std.heap.page_allocator;
- const OpType = enum(u2) {
- Sum,
- Prod,
- Softplus,
- };
- const NaiveVar = struct {
- val: f64,
- grad: f64,
- pub fn init(val: f64) NaiveVar {
- return .{ .val = val, .grad = 0.0 };
- }
- };
- const Operation = struct {
- op_type: OpType,
- inputs: [2]*NaiveVar,
- output: NaiveVar,
- };
- const NaiveTape = struct {
- ops: []Operation,
- allocator: Allocator,
- pub fn init(allocator: Allocator) NaiveTape {
- return .{
- .ops = &[_]Operation{},
- .allocator = allocator,
- };
- }
- pub fn deinit(self: *NaiveTape) void {
- self.allocator.free(self.ops);
- }
- pub fn sum(self: *NaiveTape, input1: *NaiveVar, input2: *NaiveVar) !*NaiveVar {
- const sum_val = input1.val + input2.val;
- const sum_var = NaiveVar.init(sum_val);
- const op = Operation{
- .op_type = .Sum,
- .inputs = [_]*NaiveVar{ input1, input2 },
- .output = sum_var,
- };
- self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
- self.ops[self.ops.len - 1] = op;
- return &self.ops[self.ops.len - 1].output;
- }
- pub fn prod(self: *NaiveTape, var1: *NaiveVar, var2: *NaiveVar) !*NaiveVar {
- const prod_val = var1.val * var2.val;
- const prod_var = NaiveVar.init(prod_val);
- const op = Operation{
- .op_type = .Prod,
- .inputs = [_]*NaiveVar{ var1, var2 },
- .output = prod_var,
- };
- self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
- self.ops[self.ops.len - 1] = op;
- return &self.ops[self.ops.len - 1].output;
- }
- pub fn softplus(self: *NaiveTape, nvar: *NaiveVar) !*NaiveVar {
- const softplus_val = std.math.log1p(std.math.exp(nvar.val));
- const softplus_var = NaiveVar.init(softplus_val);
- const op = Operation{
- .op_type = .Softplus,
- .inputs = [_]*NaiveVar{nvar, undefined}, // Only one input needed, second is unused
- .output = softplus_var,
- };
- self.ops = try self.allocator.realloc(self.ops, self.ops.len + 1);
- self.ops[self.ops.len - 1] = op;
- return &self.ops[self.ops.len - 1].output;
- }
- pub fn backward(self: *NaiveTape, nvar: *NaiveVar) void {
- nvar.grad += 1.0;
- var i = self.ops.len;
- while (i > 0) {
- i -= 1;
- const op = self.ops[i];
- switch (op.op_type) {
- .Sum => {
- const output_grad = op.output.grad;
- op.inputs[0].grad += output_grad;
- op.inputs[1].grad += output_grad;
- },
- .Prod => {
- const output_grad = op.output.grad;
- const input1_val = op.inputs[0].val;
- const input2_val = op.inputs[1].val;
- op.inputs[0].grad += input2_val * output_grad;
- op.inputs[1].grad += input1_val * output_grad;
- },
- .Softplus => {
- const output_grad = op.output.grad;
- const input_val = op.inputs[0].val;
- op.inputs[0].grad += output_grad / (1.0 + std.math.exp(-input_val));
- },
- }
- }
- }
- };
- pub fn main() !void {
- var arena = ArenaAlloc.init(pageAlloc);
- defer arena.deinit();
- const allocator = arena.allocator();
- const iterations: usize = 1000000;
- const start_time = time();
- var i: usize = 0;
- while (i < iterations) : (i += 1) {
- var var1 = NaiveVar.init(1.0);
- var var2 = NaiveVar.init(2.0);
- var tape = NaiveTape.init(allocator);
- defer tape.deinit();
- const sum_var = try tape.sum(&var1, &var2);
- const prod_var = try tape.prod(sum_var, sum_var);
- const softplus_var = try tape.softplus(prod_var);
- tape.backward(softplus_var);
- if (i == iterations - 1) {
- print("sum_var val: {d:}\n", .{sum_var.val});
- print("prod_var val: {d:}\n", .{prod_var.val});
- print("softplus_var val: {d:}\n", .{softplus_var.val});
- print("sum_var grad: {d:}\n", .{sum_var.grad});
- print("var1 grad: {d:}\n", .{var1.grad});
- print("var2 grad: {d:}\n", .{var2.grad});
- }
- }
- const end_time = time();
- const elapsed_time = @as(f64, @floatFromInt(end_time - start_time));
- print("\nElapsed time: {d:.3} ms\n", .{elapsed_time});
- }
Advertisement
Add Comment
Please, Sign In to add comment