为何实现RL电路电感电压计算的C++代码运行速度远慢于Python?
C++实现RL电路有限差分法计算远慢于Python/Matlab的问题分析与优化
我是编程新手,用Python、C和Matlab写了三个脚本,通过矩阵与有限差分法计算RL电路的电感电压。运行后发现Matlab耗时0.4秒,Python耗时1秒,**C耗时却超过30秒**。已经试过更换编译器、添加-O3和-fopenmp编译选项、用引用传递,但只减少了约1秒。希望找出问题原因并得到优化建议。
原代码展示
C++代码
#include <iostream> #include <vector> #include <cmath> #include <chrono> #include <Eigen/Dense> #include "matplotlibcpp.h" #include <omp.h> namespace plt = matplotlibcpp; double R = 10.0; double L = 2.0; double V0 = 5.0; double t1 = 1.0; double t2 = 2.0; double t_start = 0.0; double t_end = 3.0; double exact_voltage(double t) { if (t < t1) { return 0; } else if (t1 <= t && t < t2) { return V0 * std::exp(-(R / L) * (t - t1)); } else { return -V0 * std::exp(-(R / L) * (t - t2)); } } void solve(double dt, std::vector<double>& solved_s_times, std::vector<double>& solved_V_L) { int N = static_cast<int>((t_end - t_start) / dt); Eigen::MatrixXd A = Eigen::MatrixXd::Zero(N, N); Eigen::VectorXd B = Eigen::VectorXd::Zero(N); for (int i = 0; i < N; ++i) { A(i, i) = L / dt + R; if (i > 0) { A(i, i - 1) = -L / dt; } } std::vector<double> solve_times(N); for (int i = 0; i < N; ++i) { solve_times[i] = t_start + i * dt; if (t1 <= solve_times[i] && solve_times[i] < t2) { B(i) = V0; } else if (solve_times[i] >= t2) { B(i) = 0.0; } } Eigen::VectorXd I = A.colPivHouseholderQr().solve(B); Eigen::MatrixXd A_dI_dt = Eigen::MatrixXd::Zero(N, N); for (int i = 0; i < N - 1; ++i) { A_dI_dt(i, i) = -1.0 / dt; A_dI_dt(i, i + 1) = 1.0 / dt; } Eigen::VectorXd dI_dt = A_dI_dt * I; solved_V_L.resize(N); solved_s_times.resize(N); for (int i = 0; i < N; ++i) { solved_V_L[i] = L * dI_dt[i]; solved_s_times[i] = solve_times[i] + dt; } } int main() { std::cout << Eigen::nbThreads() << std::endl; std::vector<double> dt_values = {0.002, 0.001, 0.004}; auto start = std::chrono::high_resolution_clock::now(); // Start pomiaru czasu for (double dt : dt_values) { std::vector<double> times, V_L; solve(dt, times, V_L); plt::plot(times, V_L, {{"label", "Symulacja (dt = " + std::to_string(dt) + ")"}}); } auto end = std::chrono::high_resolution_clock::now(); // Koniec pomiaru czasu std::chrono::duration<double> duration = end - start; std::cout << "Wykonano w: " << duration.count() << "s" << std::endl; std::vector<double> exact_times, exact_V_L; double dt_exact = 0.001; int N_exact = static_cast<int>((t_end - t_start) / dt_exact); exact_times.resize(N_exact); exact_V_L.resize(N_exact); for (int i = 0; i < N_exact; ++i) { exact_times[i] = t_start + i * dt_exact; exact_V_L[i] = exact_voltage(exact_times[i]); } plt::plot(exact_times, exact_V_L, {{"label", "Dokładne rozwiązanie"}}); plt::xlabel("Czas (s)"); plt::ylabel("Napięcie (V)"); plt::title("Napięcie na cewce w obwodzie RL - symulacja i dokładne rozwiązanie"); plt::grid(true); plt::legend(); plt::show(); return 0; }
Python代码
import numpy as np import matplotlib.pyplot as plt import time from numba import njit R = 10.0 L = 2.0 V0 = 5.0 t_start = 0.0 t_end = 3 dt_values = [0.002, 0.001, 0.004] plt.figure(figsize=(12.8, 7.2)) t1 = 1 t2 = 2 # @njit def exact_voltage(t): if t < t1: return 0 elif t1 <= t < t2: return V0 * np.exp(-(R / L) * (t - t1)) else: return -V0 * np.exp(-(R / L) * (t - t2)) # @njit def solve(dt): solve_times = np.arange(t_start, t_end, dt) N = len(solve_times) A = np.zeros((N, N)) B = np.zeros(N) for i in range(N): A[i, i] = L / dt + R if i > 0: A[i, i - 1] = -L / dt for i in range(N): if t1 <= solve_times[i] < t2: B[i] = V0 elif solve_times[i] >= t2: B[i] = 0.0 I = np.linalg.solve(A, B) A_dI_dt = np.zeros((N, N)) for i in range(N - 1): A_dI_dt[i, i] = - 1 / dt A_dI_dt[i, i + 1] = 1 / dt dI_dt = np.dot(A_dI_dt, I) solved_V_L = L * dI_dt solved_s_times = np.zeros(N) for i in range(N): solved_s_times[i] = solve_times[i] + dt return solved_s_times, solved_V_L start = time.time() for dt_value in dt_values: times, V_L = solve(dt_value) plt.plot(times, V_L, label=f'Symulacja (dt = {dt_value})') end = time.time() print(f"Wykonano w: {end - start}s") exact_times = np.arange(t_start, t_end, 0.001) exact_V_L = [exact_voltage(t) for t in exact_times] plt.plot(exact_times, exact_V_L, 'k--', label='Dokładne rozwiązanie') plt.xlabel('Czas (s)') plt.ylabel('Napięcie (V)') plt.title('Napięcie na cewce w obwodzie RL - symulacja i dokładne rozwiązanie') plt.grid(True) plt.legend() plt.show()
核心问题分析
- 矩阵存储与求解器选择错误:你的C++代码用稠密矩阵存储三对角矩阵,而Python的
numpy.linalg.solve会自动识别稀疏结构并调用优化算法;Eigen的colPivHouseholderQr()是通用稠密矩阵求解器,对稀疏矩阵效率极低,这是性能差距的核心原因。 - 不必要的稠密矩阵运算:计算
dI_dt时构建了完整的N×N差分矩阵,实际只是简单的一阶差分操作,完全不需要矩阵乘法,浪费大量内存和计算资源。 - 编译配置未完全生效:虽然添加了
-O3和-fopenmp,但Eigen默认可能未启用OpenMP加速,或编译时未正确关联优化模块。
优化建议与修改代码
优化方向
- 使用Eigen稀疏矩阵模块存储三对角结构,搭配稀疏求解器
- 用向量直接计算差分,避免冗余矩阵操作
- 确保编译时启用Eigen的并行优化
修改后的C++核心代码
#include <iostream> #include <vector> #include <cmath> #include <chrono> #include <Eigen/Sparse> #include <Eigen/Dense> #include "matplotlibcpp.h" namespace plt = matplotlibcpp; double R = 10.0; double L = 2.0; double V0 = 5.0; double t1 = 1.0; double t2 = 2.0; double t_start = 0.0; double t_end = 3.0; double exact_voltage(double t) { if (t < t1) { return 0; } else if (t1 <= t && t < t2) { return V0 * std::exp(-(R / L) * (t - t1)); } else { return -V0 * std::exp(-(R / L) * (t - t2)); } } void solve(double dt, std::vector<double>& solved_s_times, std::vector<double>& solved_V_L) { int N = static_cast<int>((t_end - t_start) / dt); // 稀疏矩阵存储三对角结构 Eigen::SparseMatrix<double> A(N, N); Eigen::VectorXd B = Eigen::VectorXd::Zero(N); A.reserve(Eigen::VectorXi::Constant(N, 2)); // 每行最多2个非零元素 for (int i = 0; i < N; ++i) { A.insert(i, i) = L / dt + R; if (i > 0) { A.insert(i, i - 1) = -L / dt; } } A.makeCompressed(); std::vector<double> solve_times(N); for (int i = 0; i < N; ++i) { solve_times[i] = t_start + i * dt; if (t1 <= solve_times[i] && solve_times[i] < t2) { B(i) = V0; } else if (solve_times[i] >= t2) { B(i) = 0.0; } } // 稀疏LU求解器 Eigen::SparseLU<Eigen::SparseMatrix<double>> solver; solver.compute(A); Eigen::VectorXd I = solver.solve(B); // 直接向量运算计算差分,避免矩阵乘法 Eigen::VectorXd dI_dt(N); dI_dt.head(N-1) = (I.tail(N-1) - I.head(N-1)) / dt; dI_dt(N-1) = 0; // 最后一个点差分设为0,可按需调整 solved_V_L.resize(N); solved_s_times.resize(N); for (int i = 0; i < N; ++i) { solved_V_L[i] = L * dI_dt[i]; solved_s_times[i] = solve_times[i] + dt; } } int main() { std::vector<double> dt_values = {0.002, 0.001, 0.004}; auto start = std::chrono::high_resolution_clock::now(); for (double dt : dt_values) { std::vector<double> times, V_L; solve(dt, times, V_L); plt::plot(times, V_L, {{"label", "Symulacja (dt = " + std::to_string(dt) + ")"}}); } auto end = std::chrono::high_resolution_clock::now(); std::chrono::duration<double> duration = end - start; std::cout << "Wykonano w: " << duration.count() << "s" << std::endl; // 精确解部分保持不变 std::vector<double> exact_times, exact_V_L; double dt_exact = 0.001; int N_exact = static_cast<int>((t_end - t_start) / dt_exact); exact_times.resize(N_exact); exact_V_L.resize(N_exact); for (int i = 0; i < N_exact; ++i) { exact_times[i] = t_start + i * dt_exact; exact_V_L[i] = exact_voltage(exact_times[i]); } plt::plot(exact_times, exact_V_L, {{"label", "Dokładne rozwiązanie"}}); plt::xlabel("Czas (s)"); plt::ylabel("Napięcie (V)"); plt::title("Napięcie na cewce w obwodzie RL - symulacja i dokładne rozwiązanie"); plt::grid(true); plt::legend(); plt::show(); return 0; }
推荐编译命令
g++ -O3 -fopenmp -DEIGEN_USE_OPENMP your_code.cpp -o rl_simulation -I/path/to/eigen
内容的提问来源于stack exchange,提问作者Andrzej Sołtys
相关产品推荐
相关产品推荐

