A shape describes the number of dimensions in a array, the bounds of each dimension, and the primitive component type. For tuples, shape describes the structure (number of elements and nesting).
| 33 | // dimension, and the primitive component type. For tuples, shape describes the |
| 34 | // structure (number of elements and nesting). |
| 35 | class Shape { |
| 36 | public: |
| 37 | Shape() = default; |
| 38 | |
| 39 | // Construct a shape from a ShapeProto. |
| 40 | explicit Shape(const ShapeProto& shape_proto); |
| 41 | |
| 42 | // Returns a ShapeProto representation of the Shape. |
| 43 | ShapeProto ToProto() const; |
| 44 | |
| 45 | // Returns a human-readable string that represents the given shape, with or |
| 46 | // without layout. e.g. "F32[42,12] {0, 1}" or "F32[64]". |
| 47 | string ToString(bool print_layout = false) const; |
| 48 | |
| 49 | // Returns the rank (number of dimensions) of the given shape. Shape must be |
| 50 | // an array. |
| 51 | int64 rank() const { |
| 52 | CHECK(IsArray()) << "Non-arrays do not have a rank, shape: " << ToString(); |
| 53 | return dimensions_.size(); |
| 54 | } |
| 55 | |
| 56 | // Returns whether the shape is of the specified type (array, tuple, etc). |
| 57 | bool IsArray() const { return primitive_util::IsArrayType(element_type()); } |
| 58 | bool IsTuple() const { return element_type() == TUPLE; } |
| 59 | bool IsToken() const { return element_type() == TOKEN; } |
| 60 | bool IsOpaque() const { return element_type() == OPAQUE_TYPE; } |
| 61 | |
| 62 | // Returns true if no array dimension in the shape is dynamically sized. Tuple |
| 63 | // shapes are traversed recursively. |
| 64 | bool is_static() const; |
| 65 | |
| 66 | // Returns true if the given dimension is dynamically-sized. |
| 67 | bool is_dynamic_dimension(int dimension) const { |
| 68 | return dynamic_dimensions_.at(dimension); |
| 69 | } |
| 70 | |
| 71 | // Sets whether or not the given dimension is dynamically-sized. |
| 72 | void set_dynamic_dimension(int dimension, bool is_dynamic) { |
| 73 | dynamic_dimensions_[dimension] = is_dynamic; |
| 74 | } |
| 75 | |
| 76 | absl::Span<const bool> dynamic_dimensions() const { |
| 77 | return dynamic_dimensions_; |
| 78 | } |
| 79 | |
| 80 | absl::Span<bool> mutable_dynamic_dimensions() { |
| 81 | return absl::MakeSpan(dynamic_dimensions_); |
| 82 | } |
| 83 | |
| 84 | // Add dimension_upper_bound(). |
| 85 | |
| 86 | // Removes the given dimension form the shape. Layout, if it exists, is |
| 87 | // adjusted to match the modified shape. |
| 88 | void DeleteDimension(int64 dim_to_delete); |
| 89 | |
| 90 | // The following methods mirror the protobuf generated code interface for the |
| 91 | // message ShapeProto. This enabled easy migration of this data structure |
| 92 | // from a proto to a proper C++ class. |