载入中...
搜索中...
未找到
GpuGraph.cpp
浏览该文件的文档.
2
3namespace eve::agent::detail {
4namespace {
6struct Builder {
7 tensor::Graph g;
8 int node(OpType op, int rows, int cols, int a = -1, int b = -1) {
9 tensor::GraphNode n;
10 n.type = op;
11 n.rank = 2;
12 n.dims[0] = rows;
13 n.dims[1] = cols;
14 n.in0 = a;
15 n.in1 = b;
16 return g.addNode(std::move(n));
17 }
18 int input(int slot, int rows, int cols) {
19 int id = node(OpType::Placeholder, rows, cols);
20 g.node(id).placeholderSlot = slot;
21 return id;
22 }
23 int unary(OpType op, int a, float s0 = 0, float s1 = 0) {
24 auto shape = g.node(a);
25 int id = node(op, shape.dims[0], shape.dims[1], a);
26 g.node(id).s0 = s0;
27 g.node(id).s1 = s1;
28 return id;
29 }
30 int binary(OpType op, int a, int b) {
31 auto shape = g.node(a);
32 return node(op, shape.dims[0], shape.dims[1], a, b);
33 }
34 int transpose(int a) {
35 auto s = g.node(a);
36 int id = node(OpType::Permute, s.dims[1], s.dims[0], a);
37 g.node(id).permRank = 2;
38 g.node(id).perm[0] = 1;
39 g.node(id).perm[1] = 0;
40 return id;
41 }
42 int matmul(int a, int b) { return node(OpType::MatMul, g.node(a).dims[0], g.node(b).dims[1], a, b); }
43 int slice(int a, int axis, int begin, int end) {
44 auto s = g.node(a);
45 s.dims[axis] = end - begin;
46 int id = node(OpType::Slice, s.dims[0], s.dims[1], a);
47 g.node(id).i0 = axis;
48 g.node(id).i1 = begin;
49 g.node(id).i2 = end;
50 return id;
51 }
52 int concat(int a, int b) {
53 int id = node(OpType::Concat, g.node(a).dims[0], g.node(a).dims[1] + g.node(b).dims[1], a, b);
54 g.node(id).i0 = 1;
55 g.node(id).i1 = 2;
56 return id;
57 }
58 int biasInput(int a) {
59 int one = node(OpType::Const, 1, 1);
60 g.node(one).constData = {1};
61 return concat(a, one);
62 }
63};
64} // namespace
65PolicyGraph makePolicyGraph(int f, int h, int a, bool training) {
66 Builder b;
67 const int n1 = h * (f + 1), n2 = h * (h + 1), n3 = a * (h + 1);
68 const int weights = b.input(0, 1, n1 + n2 + n3);
69 const int x = b.input(1, 1, f), mask = b.input(2, 1, a);
70 auto matrix = [&](int start, int count, int rows, int cols) {
71 int flat = b.slice(weights, 1, start, start + count);
72 return b.node(OpType::Reshape, rows, cols, flat);
73 };
74 const int w1 = matrix(0, n1, h, f + 1), w2 = matrix(n1, n2, h, h + 1), w3 = matrix(n1 + n2, n3, a, h + 1);
75 const int xb = b.biasInput(x);
76 const int h1 = b.unary(OpType::Tanh, b.matmul(xb, b.transpose(w1)));
77 const int h1b = b.biasInput(h1);
78 const int h2 = b.unary(OpType::Tanh, b.matmul(h1b, b.transpose(w2)));
79 const int h2b = b.biasInput(h2);
80 const int logits = b.binary(OpType::Add, b.matmul(h2b, b.transpose(w3)), mask);
81 int output = b.unary(OpType::Softmax, logits);
82 b.g.node(output).i0 = 1;
83 if (training) {
84 const int target = b.input(3, 1, a), rate = b.input(4, 1, 1);
85 const int d3 = b.binary(OpType::Sub, output, target);
86 auto delta = [&](int d, int w, int activation) {
87 int propagated = b.matmul(d, b.slice(w, 1, 0, h));
88 int derivative =
89 b.unary(OpType::AddScalar, b.unary(OpType::Neg, b.binary(OpType::Multiply, activation, activation)), 1);
90 return b.binary(OpType::Multiply, propagated, derivative);
91 };
92 const int d2 = delta(d3, w3, h2), d1 = delta(d2, w2, h1);
93 auto update = [&](int w, int d, int input, int count) {
94 int gradient = b.unary(OpType::Clamp, b.matmul(b.transpose(d), input), -1, 1);
95 int updated = b.binary(OpType::Sub, w, b.binary(OpType::Multiply, gradient, rate));
96 return b.node(OpType::Reshape, 1, count, updated);
97 };
98 int u1 = update(w1, d1, xb, n1), u2 = update(w2, d2, h1b, n2), u3 = update(w3, d3, h2b, n3);
99 output = b.concat(b.concat(u1, u2), u3);
100 }
101 return {std::move(b.g), output};
102}
103} // namespace eve::agent::detail
LogicalId target
Duration start
float w
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
std::string output
const std::string & s
int mask
EvpackChunkInput input
Definition Evpack.cpp:170
int rows
int cols
tensor::Graph g
Definition GpuGraph.cpp:7
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
int h
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
float f
std::array< std::uint64_t, kPixelChunkSize *kPixelChunkSize > updated
std::string id
Definition PlayHost.cpp:108
float begin
float d
const RoadNode * node
std::uint32_t count
float weights[3]
PolicyGraph makePolicyGraph(int f, int h, int a, bool training)
Make policy graph.
Definition GpuGraph.cpp:65
Result< std::vector< int32_t > > matmul(ByteView a, ByteView b, size_t m, size_t k, size_t n, int aZero, int bZero, OnnxCompute *compute)
Integer row-major [M,K] x [K,N], subtracting scalar zero points.
void concat(const float *const *ins, const int *const *inDims, const int *inRanks, int n, int axis, float *out, const int *outDims, int outRank)
Concat.
int axis(int64_t a, size_t rank)
Axis.
OpType
OpType public API.
Definition Graph.h:22
PolicyGraph public API.
Definition GpuGraph.h:9