-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbackend_training.zig
More file actions
97 lines (78 loc) · 3.16 KB
/
Copy pathbackend_training.zig
File metadata and controls
97 lines (78 loc) · 3.16 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
const std = @import("std");
const nn = @import("nn");
const Matrix = nn.Matrix;
const BackendMatrix = nn.BackendMatrix;
const TrainingResult = struct {
backend_name: []const u8,
before_loss: f64,
after_loss: f64,
prediction_at_two: f64,
};
pub fn main() !void {
var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
defer arena.deinit();
const result = try runBackendTraining(arena.allocator());
std.debug.print("Backend training demo ({s})\n", .{result.backend_name});
std.debug.print("loss: {d:.6} -> {d:.6}\n", .{ result.before_loss, result.after_loss });
std.debug.print("prediction for x=2: {d:.4}\n", .{result.prediction_at_two});
}
fn requestedBackendType() nn.BackendType {
if (nn.enable_cuda) return .CUDA;
if (nn.enable_rocm) return .ROCm;
if (nn.enable_metal) return .Metal;
return .CPU;
}
fn backendName(backend_type: nn.BackendType) []const u8 {
return switch (backend_type) {
.CPU => "cpu",
.Metal => "metal",
.CUDA => "cuda",
.ROCm => "rocm",
};
}
fn runBackendTraining(allocator: std.mem.Allocator) !TrainingResult {
var network = nn.Network.init(allocator, 0.05, .MeanSquaredError);
defer network.deinit();
try network.addLayer(1, 1, nn.Activation.linear, nn.Activation.linear_derivative);
const layer = network.layers.items[0].Standard;
layer.weights.fill(0.0);
layer.bias.fill(0.0);
var inputs = try Matrix.init(allocator, 4, 1);
defer inputs.deinit();
var targets = try Matrix.init(allocator, 4, 1);
defer targets.deinit();
for (0..4) |row| {
const x = @as(f64, @floatFromInt(row));
try inputs.set(row, 0, x);
try targets.set(row, 0, x * 2.0);
}
var backend = try nn.createBackend(allocator, requestedBackendType());
defer backend.deinit();
const backend_inputs = try BackendMatrix.fromMatrix(backend, inputs, allocator);
defer backend_inputs.deinit();
const backend_targets = try BackendMatrix.fromMatrix(backend, targets, allocator);
defer backend_targets.deinit();
var trainer = try network.backendTrainerWithOptimizer(backend, nn.BackendOptimizerConfig.withMomentum(0.9));
defer trainer.deinit();
const before_predictions = try trainer.predict(backend_inputs);
defer before_predictions.deinit();
const before_loss = try trainer.calculateLoss(before_predictions, backend_targets);
for (0..32) |_| {
_ = try trainer.trainBatch(backend_inputs, backend_targets);
}
const after_predictions = try trainer.predict(backend_inputs);
defer after_predictions.deinit();
const after_loss = try trainer.calculateLoss(after_predictions, backend_targets);
try trainer.syncToNetwork(&network);
return .{
.backend_name = backendName(backend.getBackendType()),
.before_loss = before_loss,
.after_loss = after_loss,
.prediction_at_two = after_predictions.get(2, 0),
};
}
test "backend training experiment reduces loss" {
const result = try runBackendTraining(std.testing.allocator);
try std.testing.expect(result.after_loss < result.before_loss);
try std.testing.expect(result.prediction_at_two > 2.0);
}