载入中...
搜索中...
未找到
OnnxShapeOps.cpp
浏览该文件的文档.
18size_t broadcastIndex(size_t i, const std::vector<int64_t>& shape, const std::vector<int64_t>& output) {
31std::optional<RuntimeTensor> executeShape(const Node& n, const std::vector<const RuntimeTensor*>& in) {
38 if (starts.size() != ends.size() || starts.size() != axes.size() || starts.size() != steps.size())
45 if (!seen.insert(a).second || !steps[j] || steps[j] == INT64_MIN) throw Failure("Invalid Slice axis/step");
58 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
67 std::memcpy(out.bytes.data() + i * elementSize(x.element), x.bytes.data() + source * elementSize(x.element),
83 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
100 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
115 RuntimeTensor out{value.element, shape, std::vector<uint8_t>(count(shape) * elementSize(value.element))};
123 if (x.element != OnnxElement::Bool || a.element != b.element) throw Failure("Where type mismatch");
129 std::memcpy(out.bytes.data() + i * size, v.bytes.data() + broadcastIndex(i, v.shape, shape) * size, size);
153 out.bytes[i] = x.element == OnnxElement::Float32 ? compare(read<float>(x, a), read<float>(y, b))
158 if ((n.op == "Add" || n.op == "Sub" || n.op == "Mul" || n.op == "Div") && x.element != OnnxElement::Float32) {
160 if (x.element != y.element || (x.element != OnnxElement::Int64 && x.element != OnnxElement::Int32))
163 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
Definition OnnxByteStorage.h:6
size_t broadcastIndex(size_t i, const std::vector< int64_t > &shape, const std::vector< int64_t > &output)
Broadcast index.
Definition OnnxShapeOps.cpp:18
RuntimeTensor make(OnnxElement e, std::vector< int64_t > shape, const std::vector< T > &values)
Make.
Definition OnnxInternal.h:91
std::vector< int64_t > broadcast(const std::vector< int64_t > &a, const std::vector< int64_t > &b)
Broadcast.
Definition OnnxShapeOps.cpp:8
std::optional< RuntimeTensor > executeShape(const Node &, const std::vector< const RuntimeTensor * > &)
Execute shape.
Definition OnnxShapeOps.cpp:31
@ Float32