lib/accy/src/choir/activation.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 pub const Kind = enum {
4 gelu,
5 relu,
6 silu,
7
8 pub fn fromName(name: []const u8) ?Kind {
9 inline for (
10 @typeInfo(Kind).@"enum".field_names,
11 @typeInfo(Kind).@"enum".field_values,
12 ) |field_name, field_name_value| {
13 const field = .{ .name = field_name, .value = field_name_value };
14 if (std.mem.eql(u8, name, field.name)) return @fromBackingInt(@intCast(field.value));
15 }
16 return null;
17 }
18 };
19
20 test "activation kind parses canonical tags" {
21 try std.testing.expectEqual(@as(?Kind, .gelu), Kind.fromName("gelu"));
22 try std.testing.expectEqual(@as(?Kind, .relu), Kind.fromName("relu"));
23 try std.testing.expectEqual(@as(?Kind, .silu), Kind.fromName("silu"));
24 try std.testing.expectEqual(@as(?Kind, null), Kind.fromName("swiglu"));
25 }