diff --git a/backends/native/runtime/graph/Ids.h b/backends/native/runtime/graph/Ids.h new file mode 100644 index 00000000000..baa33415006 --- /dev/null +++ b/backends/native/runtime/graph/Ids.h @@ -0,0 +1,33 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#pragma once + +#include +#include +#include + +namespace ptn { + +// Index-arena handles: a NodeId indexes the graph's node arena, a ValueId +// its value arena. Plain int32_t aliases — they index, compare, and hash +// directly, at the cost of no NodeId/ValueId type distinction. kInvalid marks +// "no id". +using NodeId = int32_t; +using ValueId = int32_t; +inline constexpr int32_t kInvalid = -1; + +constexpr bool valid(int32_t id) { + return id >= 0; +} + +// std::cmp_less compares the signed id against the unsigned size without +// casting either side. +constexpr bool in_bounds(int32_t id, size_t size) { + return valid(id) && std::cmp_less(id, size); +} + +} // namespace ptn diff --git a/backends/native/runtime/graph/Value.cpp b/backends/native/runtime/graph/Value.cpp new file mode 100644 index 00000000000..17a624eea65 --- /dev/null +++ b/backends/native/runtime/graph/Value.cpp @@ -0,0 +1,37 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#include + +#include + +namespace ptn { + +const TensorMeta& Value::tensor_meta() const { + const TensorMeta* m = std::get_if(&value_); + if (m == nullptr) { + throw std::runtime_error("Value::tensor_meta: value is not a Tensor"); + } + return *m; +} + +const Scalar& Value::scalar() const { + const Scalar* s = std::get_if(&value_); + if (s == nullptr) { + throw std::runtime_error("Value::scalar: value is not a Scalar"); + } + return *s; +} + +const std::vector& Value::content_ids() const { + const std::vector* ids = std::get_if>(&value_); + if (ids == nullptr) { + throw std::runtime_error("Value::content_ids: value is not a List"); + } + return *ids; +} + +} // namespace ptn diff --git a/backends/native/runtime/graph/Value.h b/backends/native/runtime/graph/Value.h new file mode 100644 index 00000000000..36ff4b2efaf --- /dev/null +++ b/backends/native/runtime/graph/Value.h @@ -0,0 +1,98 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +namespace ptn { + +enum class ValueKind : int8_t { + None = 0, + Tensor = 1, + Scalar = 2, + List = 3, +}; + +// A single SSA value (dataflow edge) in a Graph: its contents plus def-use +// wiring, a storage alias and an open annotation map. The id fields are plain +// handles; whether one is in range is a property of the owning arena, so +// nothing here validates them. +// +// The variant's alternatives are listed in ValueKind order, so kind() is its +// index. A Tensor carries metadata only, so a weight is an ordinary arena +// value like any other, with its bytes held outside the graph. A List holds +// ValueIds to its elements, so nesting goes through the arena; nothing +// deserialized is a List, it exists for in-memory rewrites such as grouping a +// tuple. +class Value { + private: + std::variant> value_; + + public: + // SSA name, scoped to the enclosing Graph. + std::string name; + // Defining node; invalid => graph input. + NodeId producer_id = kInvalid; + // Def-use, built by inverting node inputs. + std::vector consumer_ids; + // Shares storage with this value (a view); fresh if invalid. + ValueId alias_id = kInvalid; + // Open annotations for graph passes and engines, like node.meta in FX. + std::unordered_map attrs; + + Value() = default; // a None value with an empty name + + explicit Value(std::string name) // a named None value + : name(std::move(name)) {} + + Value(std::string name, TensorMeta meta) + : value_(std::move(meta)), name(std::move(name)) {} + + // The empty dim_order_hint is what makes the tensor contiguous. + Value(std::string name, ScalarType dtype, std::vector sizes) + : value_(TensorMeta{dtype, std::move(sizes), {}}), + name(std::move(name)) {} + + Value(std::string name, Scalar value) + : value_(value), name(std::move(name)) {} + + Value(std::string name, std::vector elem_ids) + : value_(std::move(elem_ids)), name(std::move(name)) {} + + ValueKind kind() const { + return static_cast(value_.index()); + } + bool is_tensor() const { + return std::holds_alternative(value_); + } + bool is_scalar() const { + return std::holds_alternative(value_); + } + bool is_list() const { + return std::holds_alternative>(value_); + } + bool is_none() const { + return std::holds_alternative(value_); + } + + // Typed payload accessors: throw std::runtime_error unless the kind matches. + const TensorMeta& tensor_meta() const; + const Scalar& scalar() const; + const std::vector& content_ids() const; +}; + +} // namespace ptn diff --git a/backends/native/runtime/graph/targets.bzl b/backends/native/runtime/graph/targets.bzl index 7608bc89553..c18b4f522b1 100644 --- a/backends/native/runtime/graph/targets.bzl +++ b/backends/native/runtime/graph/targets.bzl @@ -38,6 +38,28 @@ def define_common_targets(): visibility = ["//executorch/backends/native/..."], ) + runtime.cxx_library( + name = "ids", + exported_headers = [ + "Ids.h", + ], + visibility = ["//executorch/backends/native/..."], + ) + + runtime.cxx_library( + name = "value", + srcs = ["Value.cpp"], + exported_headers = [ + "Value.h", + ], + exported_deps = [ + ":ids", + ":scalar", + ":tensor_meta", + ], + visibility = ["//executorch/backends/native/..."], + ) + # utils/ has no BUCK of its own, so the IR printer's target lives here. Kept # separate from the IR libraries so only a consumer that dumps the IR links # the formatting code.