如何将形状为(3,2,2)的NumPy数组重塑为(2,2,3)适配matplotlib?
Hey there! Let's work through this reshaping problem so your array plays nicely with matplotlib.pyplot.imshow.
The Problem Breakdown
Your original array a has shape (3, 2, 2) — think of this as 3 separate 2x2 matrices, one for each RGB channel. What you need is a (2, 2, 3) array, where every pixel position (row, column) holds the corresponding values from all three channels.
Why reshape() Isn't Cutting It
np.reshape() only rearranges elements in memory, without respecting the logical grouping of your axes. Whether you use the default row-major (order='C') or column-major (order='F') order, it won't group the channel values for each pixel together. That's why you're not getting the result you want.
The Correct Solutions
You need to reorder the axes of your array instead of just reshaping it. Here are two straightforward ways to do this:
1. Use np.transpose()
This function lets you explicitly define the new order of axes. Your original axes are 0 (channels), 1 (rows), 2 (columns). We want to reorder them to 1 (rows), 2 (columns), 0 (channels):
import numpy as np # Example input array a = np.array([[[a1,a2],[a3,a4]], [[b1,b2],[b3,b4]], [[c1,c2],[c3,c4]]]) # Reorder axes to (rows, columns, channels) target_array = a.transpose(1, 2, 0)
Now target_array will have shape (2, 2, 3) and match your desired output exactly:
[[[a1,b1,c1], [a2,b2,c2]], [[a3,b3,c3], [a4,b4,c4]]]
2. Use np.moveaxis() (More Intuitive)
If you just need to shift the channel axis from the first position to the last, moveaxis() is even simpler. It takes the axis you want to move and its new position:
target_array = np.moveaxis(a, source=0, destination=-1)
This does the exact same thing as the transpose example but is easier to read — we're moving the 0th axis (channels) to the final position (-1 refers to the last axis in the array).
Verifying and Using with imshow
You can confirm the shape and values with:
print(target_array.shape) # Should output (2, 2, 3) print(target_array[0, 0]) # Should output [a1, b1, c1]
Since your array already contains float values in the [0, 1] range, you can pass it directly to imshow:
import matplotlib.pyplot as plt plt.imshow(target_array) plt.show()
内容的提问来源于stack exchange,提问作者Jon

