如何删除含特定值的三维数组行及对应预期数组行(Numpy实现)
Hey there! Let's work through this problem together. I get that you're trying to strip out rows containing the value 1.4567 from your 3D NumPy arrays actual and expected, while making sure their corresponding rows stay in sync. Your earlier attempts with loops or the axis parameter didn't pan out, so here's the proper NumPy way to handle this—no messy loops required.
First, let's clarify the approach
The core idea is to create a mask that identifies which rows to keep (i.e., rows that don't contain 1.4567), then apply that mask to both arrays to filter them consistently. We'll cover two common scenarios depending on whether you want to keep the original 3D grouping structure or flatten the results into a 2D array.
Scenario 1: Keep the original 3D grouping structure
If your arrays are shaped like (num_groups, num_rows, num_features) (where each group is a 2D sub-array), and you want to filter rows within each group while preserving the group structure:
import numpy as np # Example 3D arrays actual = np.array([ [[1.0, 2.0], [1.4567, 3.0], [4.0, 5.0]], [[6.0, 7.0], [8.0, 9.0], [10.0, 1.4567]] ]) expected = np.array([ [[0, 1], [2, 3], [4, 5]], [[6, 7], [8, 9], [10, 11]] ]) # Filter rows within each group filtered_actual = np.array([ group[~np.isclose(group, 1.4567).any(axis=1)] for group in actual ]) filtered_expected = np.array([ group[~np.isclose(group, 1.4567).any(axis=1)] for group in expected ])
Let's break this down:
np.isclose(group, 1.4567): Safely checks for the target value (better than==for floating-point numbers to avoid precision issues)..any(axis=1): Checks if any element in a row matches the target value.~: Flips the boolean result to get rows we want to keep (instead of delete).- The list comprehension processes each group independently, then converts back to a NumPy array.
Scenario 2: Flatten into a 2D array (ignore original groups)
If you don't need to preserve the group structure and just want all valid rows in a single 2D array:
# Flatten both arrays to 2D (all rows combined) actual_2d = actual.reshape(-1, actual.shape[-1]) expected_2d = expected.reshape(-1, expected.shape[-1]) # Create a mask for rows without 1.4567 mask = ~np.isclose(actual_2d, 1.4567).any(axis=1) # Apply the mask to both arrays filtered_actual = actual_2d[mask] filtered_expected = expected_2d[mask]
This approach simplifies things by treating all rows as a single collection, then filtering them in one go.
Why your earlier attempts might have failed
- If you tried using
np.deletewithaxis=1, NumPy struggles with this when different groups have different numbers of rows to delete (it requires uniform dimensions across the array). The list comprehension workaround handles variable row counts per group perfectly. - Using loops without vectorization can lead to slow code or off-by-one errors—NumPy's vectorized operations are faster and more reliable here.
内容的提问来源于stack exchange,提问作者Lucy Brown

