【发布时间】:2017-07-02 21:13:21
【问题描述】:
我正在考虑从 Matlab 切换到 Python (NumPy)。因此,作为一项学习任务,我尝试在 Python 上重写一个简单的随机模型。 python 脚本返回正确答案,但运行速度太慢! Python 需要 3 分钟,而不是像 Matlab 一样需要 3 秒。我做错了什么?
Matlab:
clear all; clc;
tic
T = 0.05;
Tmax = 3600;
t = T:T:Tmax;
N = length(t);
G = [0 0;
0 T];
F = [1 T;
0 1];
Dksi = 13*1;
Deta = 10*1;
Band = 0.1:0.1:3;
RMS_Omega = nan(1, length(Band));
for i = 1:length(Band)
K = nan(2, 1);
K(1) = 8/3 * Band(i) * T;
K(2) = 32/9 * Band(i)^2 * T;
ksi = sqrt(Dksi) * randn(1, N);
eta = sqrt(Deta) * randn(1, N);
Xest = [0; 0];
Xextr = F*Xest;
Xist = [0; 0];
ErrOmega = nan(1, N); Omega = nan(1, N);
for k = 1:N
Xist = F*Xist + G*[0; ksi(k)];
omega_meas = Xist(1) + eta(k);
Xest = Xextr + K*(omega_meas - Xextr(1));
Xextr = F*Xest;
ErrOmega(k) = Xest(1) - Xist(1);
Omega(k) = Xist(1);
end
RMS_Omega(i) = sqrt(mean(ErrOmega.^2));
end
figure(1)
hold on
plot(Band, RMS_Omega);
hold off
xlabel('Bandwidth, Hz'); ylabel('RMS \omega, Hz');
toc
Python:
#!/usr/bin/python
# -*- coding: utf-8 -*-
import math
import numpy as np
import matplotlib as mpl
import matplotlib.pyplot as plt
import time as time
tbeg = time.time()
T = 0.005
Tmax = 3600.0
t = np.linspace(T, Tmax, int(Tmax/T))
N = len(t)
G = np.array([[0, 0],
[0, T]])
F = np.array([[1, T],
[0, 1]])
Dksi = 13.0
Deta = 10.0
Band = np.linspace(0.1, 3.0, 30)
Band_for_plot = 2
RMS_Omega = np.array([None for i in range(0, len(Band))])
for i, BW in enumerate(Band):
K = np.array([[8/3 * BW * T],
[32/9 * BW*BW *T]])
ksi = math.sqrt(Dksi) * np.random.randn(N)
eta = math.sqrt(Deta) * np.random.randn(N)
Xest = np.array([[0],
[0]])
Xextr = F.dot(Xest)
Xist = np.array([[0],
[0]])
ErrOmega = np.array([None for j in range(0, N)])
Omega = np.array([None for j in range(0, N)])
for k in range(0, N):
Xist = F.dot(Xist) + G.dot(np.array([[0],
[ksi[k]]]))
omega_meas = Xist[0] + eta[k]
Xest = Xextr + K * (omega_meas - Xextr[0])
Xextr = F.dot(Xest)
ErrOmega[k] = Xest[0] - Xist[0]
Omega[k] = Xist[0]
RMS_Omega[i] = math.sqrt(np.mean(ErrOmega**2))
elapsed = time.time() - tbeg
print(elapsed, u'sec')
【问题讨论】:
-
分析你的python脚本会告诉你瓶颈。
-
你有巨大的原生 python 循环。那些会很慢(相对而言)。 Numpy 不是魔法仙尘,只有矢量化它才会加速你的代码。
-
是的 2 个大型普通 python 循环,这可能是你速度慢的地方,多使用 numpy
-
好吧,一方面,在 Matlab
T = 0.05中,但在 PythonT = 0.005中 - 你正在让 Python 处理 10 倍的数据。 -
您的 MATLAB 在旧版本上也会很慢。当前的 MATLAB 有一些 JIT 编译,可以让你摆脱循环。