@@ -289,6 +289,50 @@ void runGetDimTest(slim_c10::DeviceType device_type) {
289289 }
290290}
291291
292+ void runGetNumelTest (slim_c10::DeviceType device_type) {
293+ slim_c10::Device device (device_type, 0 );
294+
295+ // Test 0D tensor (scalar)
296+ {
297+ std::vector<int64_t > sizes = {};
298+ std::vector<int64_t > strides = {};
299+
300+ Tensor* tensor = new Tensor (slim::empty_strided (
301+ slim::makeArrayRef (sizes),
302+ slim::makeArrayRef (strides),
303+ slim_c10::ScalarType::Float,
304+ device));
305+
306+ int64_t ret_numel = -1 ;
307+ AOTITorchError error = aoti_torch_get_numel (tensor, &ret_numel);
308+
309+ EXPECT_EQ (error, Error::Ok);
310+ EXPECT_EQ (ret_numel, 1 );
311+
312+ delete tensor;
313+ }
314+
315+ // Test 3D tensor
316+ {
317+ std::vector<int64_t > sizes = {2 , 3 , 4 };
318+ std::vector<int64_t > strides = calculateContiguousStrides (sizes);
319+
320+ Tensor* tensor = new Tensor (slim::empty_strided (
321+ slim::makeArrayRef (sizes),
322+ slim::makeArrayRef (strides),
323+ slim_c10::ScalarType::Float,
324+ device));
325+
326+ int64_t ret_numel = -1 ;
327+ AOTITorchError error = aoti_torch_get_numel (tensor, &ret_numel);
328+
329+ EXPECT_EQ (error, Error::Ok);
330+ EXPECT_EQ (ret_numel, 24 );
331+
332+ delete tensor;
333+ }
334+ }
335+
292336// ============================================================================
293337// Storage & Device Property Tests
294338// ============================================================================
@@ -400,6 +444,10 @@ TEST_F(CommonShimsSlimTest, GetDim_CPU) {
400444 runGetDimTest (slim_c10::DeviceType::CPU );
401445}
402446
447+ TEST_F (CommonShimsSlimTest, GetNumel_CPU) {
448+ runGetNumelTest (slim_c10::DeviceType::CPU );
449+ }
450+
403451TEST_F (CommonShimsSlimTest, GetStorageOffset_CPU) {
404452 runGetStorageOffsetTest (slim_c10::DeviceType::CPU );
405453}
@@ -456,6 +504,13 @@ TEST_F(CommonShimsSlimTest, GetDim_CUDA) {
456504 runGetDimTest (slim_c10::DeviceType::CUDA );
457505}
458506
507+ TEST_F (CommonShimsSlimTest, GetNumel_CUDA) {
508+ if (!isCudaAvailable ()) {
509+ GTEST_SKIP () << " CUDA not available" ;
510+ }
511+ runGetNumelTest (slim_c10::DeviceType::CUDA );
512+ }
513+
459514TEST_F (CommonShimsSlimTest, GetStorageOffset_CUDA) {
460515 if (!isCudaAvailable ()) {
461516 GTEST_SKIP () << " CUDA not available" ;
@@ -495,13 +550,15 @@ TEST_F(CommonShimsSlimTest, NullTensorArgument) {
495550 int64_t * strides = nullptr ;
496551 int32_t dtype = -1 ;
497552 int64_t dim = -1 ;
553+ int64_t numel = -1 ;
498554
499555 EXPECT_EQ (
500556 aoti_torch_get_data_ptr (nullptr , &data_ptr), Error::InvalidArgument);
501557 EXPECT_EQ (aoti_torch_get_sizes (nullptr , &sizes), Error::InvalidArgument);
502558 EXPECT_EQ (aoti_torch_get_strides (nullptr , &strides), Error::InvalidArgument);
503559 EXPECT_EQ (aoti_torch_get_dtype (nullptr , &dtype), Error::InvalidArgument);
504560 EXPECT_EQ (aoti_torch_get_dim (nullptr , &dim), Error::InvalidArgument);
561+ EXPECT_EQ (aoti_torch_get_numel (nullptr , &numel), Error::InvalidArgument);
505562}
506563
507564TEST_F (CommonShimsSlimTest, NullReturnPointer) {
@@ -512,6 +569,7 @@ TEST_F(CommonShimsSlimTest, NullReturnPointer) {
512569 EXPECT_EQ (aoti_torch_get_strides (tensor, nullptr ), Error::InvalidArgument);
513570 EXPECT_EQ (aoti_torch_get_dtype (tensor, nullptr ), Error::InvalidArgument);
514571 EXPECT_EQ (aoti_torch_get_dim (tensor, nullptr ), Error::InvalidArgument);
572+ EXPECT_EQ (aoti_torch_get_numel (tensor, nullptr ), Error::InvalidArgument);
515573}
516574
517575// ============================================================================
@@ -534,6 +592,7 @@ TEST_F(CommonShimsSlimTest, ScalarTensor) {
534592 int64_t * ret_sizes = nullptr ;
535593 int64_t * ret_strides = nullptr ;
536594 int64_t ret_dim = -1 ;
595+ int64_t ret_numel = -1 ;
537596
538597 EXPECT_EQ (aoti_torch_get_sizes (tensor, &ret_sizes), Error::Ok);
539598 EXPECT_NE (ret_sizes, nullptr );
@@ -543,6 +602,9 @@ TEST_F(CommonShimsSlimTest, ScalarTensor) {
543602
544603 EXPECT_EQ (aoti_torch_get_dim (tensor, &ret_dim), Error::Ok);
545604 EXPECT_EQ (ret_dim, 0 );
605+
606+ EXPECT_EQ (aoti_torch_get_numel (tensor, &ret_numel), Error::Ok);
607+ EXPECT_EQ (ret_numel, 1 );
546608}
547609
548610TEST_F (CommonShimsSlimTest, LargeTensor) {
@@ -559,6 +621,7 @@ TEST_F(CommonShimsSlimTest, LargeTensor) {
559621
560622 int64_t * ret_sizes = nullptr ;
561623 int64_t * ret_strides = nullptr ;
624+ int64_t ret_numel = -1 ;
562625
563626 EXPECT_EQ (aoti_torch_get_sizes (tensor, &ret_sizes), Error::Ok);
564627 EXPECT_EQ (ret_sizes[0 ], 100 );
@@ -569,6 +632,9 @@ TEST_F(CommonShimsSlimTest, LargeTensor) {
569632 EXPECT_EQ (ret_strides[0 ], 60000 ); // 200 * 300
570633 EXPECT_EQ (ret_strides[1 ], 300 ); // 300
571634 EXPECT_EQ (ret_strides[2 ], 1 );
635+
636+ EXPECT_EQ (aoti_torch_get_numel (tensor, &ret_numel), Error::Ok);
637+ EXPECT_EQ (ret_numel, 6000000 );
572638}
573639
574640TEST_F (CommonShimsSlimTest, ConsistentPointerReturn) {
0 commit comments