This is an automated email from the ASF dual-hosted git repository.
alamb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-rs.git
The following commit(s) were added to refs/heads/main by this push:
new 4cb8ff5a49 feat(arrow-schema): expose FixedShapeTensor shape (#11247)
4cb8ff5a49 is described below
commit 4cb8ff5a4998a9fde0755fc69744a7a0d8cb8aa7
Author: yongster <[email protected]>
AuthorDate: Wed Sep 30 21:01:24 2026 +0800
feat(arrow-schema): expose FixedShapeTensor shape (#11247)
# Which issue does this PR close?
- Closes #11246.
# Rationale for this change
`FixedShapeTensor::try_new` stores the physical shape, but the type only
exposes `dimensions()` and `list_size()`. Those are the same for `[2,
6]` and `[3, 4]`, so a caller cannot recover the shape without parsing
the extension metadata JSON.
PyArrow's `FixedShapeTensorType.shape` and Arrow C++'s
`FixedShapeTensorType::shape()` already return this value.
# What changes are included in this PR?
- `FixedShapeTensor::shape(&self) -> &[usize]`
- `FixedShapeTensorMetadata::shape(&self) -> &[usize]`
- A doc example for `[2, 6]`, and a test that `[2, 6]` and `[3, 4]` stay
distinct while sharing `dimensions()` and `list_size()`
Serialization and validation are unchanged. This does not rename the
metadata key `permutations`.
# Are these changes tested?
Yes.
- `cargo test -p arrow-schema --features canonical_extension_types --lib
extension::canonical::fixed_shape_tensor`
- `cargo test -p arrow-schema --doc --features canonical_extension_types
shape`
- `cargo clippy -p arrow-schema --features canonical_extension_types
--all-targets --no-deps -- -D warnings`
# Are there any user-facing changes?
Additive API only, behind the existing `canonical_extension_types`
feature. No breaking change.
AI assistance: the accessors and tests were drafted with AI and then
checked against the existing `dimension_names()` / `permutations()`
style.
Co-authored-by: yang3.xie <[email protected]>
---
.../src/extension/canonical/fixed_shape_tensor.rs | 38 +++++++++++++++++++++-
1 file changed, 37 insertions(+), 1 deletion(-)
diff --git a/arrow-schema/src/extension/canonical/fixed_shape_tensor.rs
b/arrow-schema/src/extension/canonical/fixed_shape_tensor.rs
index d30df26088..bf85d53fe5 100644
--- a/arrow-schema/src/extension/canonical/fixed_shape_tensor.rs
+++ b/arrow-schema/src/extension/canonical/fixed_shape_tensor.rs
@@ -113,6 +113,23 @@ impl FixedShapeTensor {
self.metadata.list_size()
}
+ /// Returns the physical shape of the contained tensors.
+ ///
+ /// [`Self::dimensions`] and [`Self::list_size`] do not identify this
+ /// shape. `[2, 6]` and `[3, 4]` are both 2-dimensional and both have
+ /// 12 elements.
+ ///
+ /// ```
+ /// # use arrow_schema::extension::FixedShapeTensor;
+ /// # use arrow_schema::DataType;
+ /// let tensor = FixedShapeTensor::try_new(DataType::Float32, [2, 6],
None, None).unwrap();
+ /// assert_eq!(tensor.shape(), &[2, 6]);
+ /// assert_eq!(tensor.list_size(), 12);
+ /// ```
+ pub fn shape(&self) -> &[usize] {
+ self.metadata.shape()
+ }
+
/// Returns the number of dimensions in this fixed shape tensor.
pub fn dimensions(&self) -> usize {
self.metadata.dimensions()
@@ -334,6 +351,11 @@ impl FixedShapeTensorMetadata {
})
}
+ /// Returns the physical shape of the contained tensors.
+ pub fn shape(&self) -> &[usize] {
+ &self.shape
+ }
+
/// Returns the product of all the elements in tensor shape.
pub fn list_size(&self) -> usize {
self.shape.iter().product()
@@ -438,7 +460,7 @@ mod tests {
use crate::extension::CanonicalExtensionType;
use crate::{
Field,
- extension::{EXTENSION_TYPE_METADATA_KEY, EXTENSION_TYPE_NAME_KEY},
+ extension::{EXTENSION_TYPE_METADATA_KEY, EXTENSION_TYPE_NAME_KEY,
ExtensionType},
};
use super::*;
@@ -462,6 +484,8 @@ mod tests {
field.try_extension_type::<FixedShapeTensor>()?,
fixed_shape_tensor
);
+ assert_eq!(fixed_shape_tensor.shape(), &[100, 200, 500]);
+ assert_eq!(fixed_shape_tensor.metadata().shape(), &[100, 200, 500]);
#[cfg(feature = "canonical_extension_types")]
assert_eq!(
field.try_canonical_extension_type()?,
@@ -470,6 +494,18 @@ mod tests {
Ok(())
}
+ #[test]
+ fn shape_distinguishes_equal_sizes() -> Result<(), ArrowError> {
+ let wide = FixedShapeTensor::try_new(DataType::Float32, [2, 6], None,
None)?;
+ let tall = FixedShapeTensor::try_new(DataType::Float32, [3, 4], None,
None)?;
+ assert_eq!(wide.dimensions(), tall.dimensions());
+ assert_eq!(wide.list_size(), tall.list_size());
+ assert_eq!(wide.shape(), &[2, 6]);
+ assert_eq!(tall.shape(), &[3, 4]);
+ assert_ne!(wide.shape(), tall.shape());
+ Ok(())
+ }
+
#[test]
#[should_panic(expected = "Extension type name missing")]
fn missing_name() {