Numpy数组形状异常求助:得到(6,1,2)而非预期(6,2)
问题分析与解决方案
作为从MATLAB转Numpy的开发者,你遇到的其实是两种工具在数组维度处理上的核心差异问题,咱们一步步拆解:
为什么pos的形状是(6,1,2)而不是预期的(6,2)?
你创建pos的代码是:
pos = np.array([[ship_lats], [ship_longs]], dtype = "d").T
问题出在[[ship_lats], [ship_longs]]这层嵌套上:
ship_lats和ship_longs是Numpy一维数组,形状为(6,)——这是Numpy特有的「秩1数组」,没有行/列的概念,和MATLAB里默认的二维列向量完全不同。- 当你给每个一维数组额外套一层列表
[...]时,Numpy会把这个结构解析成三维数组,形状为(2, 1, 6)。 - 转置
.T会反转维度顺序,最终得到(6, 1, 2)的三维数组,这就是你索引pos[i][1]报错的原因——pos[i]是一个(1,2)的二维数组,它的轴0长度只有1,自然找不到索引1的元素。
而你疑惑的(6,)和(6,1)的差异:
(6,):一维数组,只是一串元素的集合,没有行/列属性,转置也不会改变形状。(6,1):二维数组,等价于MATLAB里的列向量,是6行1列的矩阵结构。
怎么修正得到(6,2)的pos数组?
有两种简洁且符合Numpy习惯的方法:
方法1:用np.column_stack(最直观)
这个函数专门用来把多个一维数组按列拼接成二维数组,直接得到你想要的形状:
pos = np.column_stack((ship_lats, ship_longs))
方法2:去掉多余嵌套后转置
删掉创建数组时的内层括号,让Numpy直接把两个一维数组拼成二维数组,再转置:
pos = np.array([ship_lats, ship_longs]).T
[ship_lats, ship_longs]会被解析成(2,6)的二维数组,转置后就是(6,2)。
修正后,你就可以正常用pos[i][1](或者更高效的Numpy写法pos[i,1])来索引经度了。
给MATLAB背景开发者的Numpy实用技巧
从MATLAB转Numpy最容易踩维度的坑,这里给你几个关键提示:
- 放弃「所有数组都是二维」的思维:Numpy支持一维数组
(n,),如果需要和MATLAB的列/行向量对应,手动用reshape转换,比如ship_lats.reshape(-1,1)就能把(6,)转成(6,1)的二维列向量。 - 数组拼接用专用函数:
- MATLAB的
[A, B](水平拼接)对应Numpy的np.hstack((A,B))或np.column_stack(针对一维数组) - MATLAB的
[A; B](垂直拼接)对应Numpy的np.vstack((A,B))或np.row_stack
- MATLAB的
- 索引习惯:Numpy支持
arr[i,j]的二维索引方式,比arr[i][j]更高效,也更接近MATLAB的arr(i,j)写法。 - 转置的区别:一维数组的
.T不会有任何变化,因为它没有第二个维度;如果要转置一维数组,先转成二维再操作,比如arr.reshape(-1,1).T。
修正后的完整代码
import pandas as pd import numpy as np import matplotlib.pyplot as plt as plot import cartopy.crs as crs import cartopy.feature as cf irmadata = pd.read_excel("DangerZone.xlsx") irma_lats = irmadata["Average Latitude"].tolist() irma_longs = irmadata["Average Longitude"].tolist() shipdata = pd.read_excel("ShipPositions.xlsx") ship_lats = shipdata["Latitude"].to_numpy() ship_longs = shipdata["Longitude"].to_numpy() # 选择一种方式创建pos pos = np.column_stack((ship_lats, ship_longs)) # 或者 pos = np.array([ship_lats, ship_longs]).T extent = [-10, -90, 0, 50] ax = plot.axes(projection = crs.PlateCarree()) ax.stock_img() ax.add_feature(cf.COASTLINE) ax.coastlines(resolution = "50m") ax.set_title("Base Map") ax.set_extent(extent) ax.plot(irma_longs, irma_lats) for i in range(len(ship_lats)): lat = pos[i, 0] lon = pos[i, 1] ax.plot(lon, lat, 'o', label = f"Ship {i+1}") # 用f-string更简洁 plot.legend() plot.show()
内容的提问来源于stack exchange,提问作者Scott Hunter
相关产品推荐
相关产品推荐

