You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用dataclass对象导致Python代码运行缓慢的优化方案咨询

博士项目Python代码优化需求

我正在为博士项目优化Python代码,需要测试不同参数的函数。由于代码中变量数量较多,我创建了4个基于@dataclass的类来分组管理变量,但将dataclass对象作为函数输入时,函数运行速度明显变慢。我尝试用Numba进行加速,但部分类型不被支持;考虑在函数内使用字典来管理变量,但不确定Numba是否支持该方式。我需要一种既能兼顾变量可读性、管理性,又能提升代码运行性能的优化方案。


测试代码

import time
import numpy as np
from numba import njit, types
from numba.experimental import jitclass
from typing import Tuple
from dataclasses import dataclass

@dataclass()
class Kpi():
    name: str
    value1: int
    value2: float

spec = [
    ('name', types.unicode_type),
    ('value1', types.float32),
    ('value2', types.int32)
]

@jitclass(spec)
class Kpi_jit:
    def __init__(self, name, value1, value2):
        self.name = name
        self.value1 = value1
        self.value2 = value2

def calc_raw_1(x1, x2, stext=None):
    if stext == 'AS':
        return (x1*x2)+(x1*x2), x1+x2+x1+x2
    else:
        return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2

@njit
def calc_accel_1(x1, x2, stext=None):
    if stext == 'AS':
        return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2
    else:
        return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2

@njit
def calc_accel_2(x1:float, x2:float, stext:str=None):
    if stext == 'AS':
        return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2
    else:
        return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2

@jit(nopython=True)
def calc_accel_3(x1:float, x2:float, stext:str=None) ->Tuple[float, float]:
    if stext == 'AS':
        return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2
    else:
        return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2

@jit(nopython=True)
def calc_accel_4(kpi) ->Tuple[float, float]:
    if kpi.name == 'AS':
        return (kpi.value1 * kpi.value2) + (kpi.value1 * kpi.value2), kpi.value1 + kpi.value2 + kpi.value1 + kpi.value2
    else:
        return (kpi.value1 * kpi.value2) * 2 + (kpi.value1 * kpi.value2), (kpi.value1 + kpi.value2 + kpi.value1 + kpi.value2) * 2

@jit(nopython=True)
def calc_accel_5(x1:float, x2:float, stext:str=None) ->Tuple[float, float]:
    if stext == 'AS':
        return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2
    else:
        return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2

x1 = 5.2
x2 = 989
filter_text = 'AB'

kpi = Kpi(name = filter_text, value1 = x1, value2 = x2)
kpi_jit = Kpi_jit(name = filter_text, value1 = x1, value2 = x2)

print("====================================================")

s = time.time()
for i in range(10000000):
    if filter_text == 'AS':
        temp1 = (x1*x2)+(x1*x2)
        temp2 = x1+x2+x1+x2
    else:
        temp1 = (x1 * x2)*2 + (x1 * x2)
        temp2 = (x1 + x2 + x1 + x2) *2
print(f"calc in fnc : {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_raw_1(x1, x2, filter_text)
print(f"calc_raw_1  : {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_1(x1, x2, filter_text)
print(f"calc_accel_1: {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_2(x1, x2, filter_text)
print(f"calc_accel_2: {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_3(x1, x2, filter_text)
print(f"calc_accel_3: {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_4(kpi_jit)
print(f"calc_accel_4: {round(time.time()- s, 2)} sec.")

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_5(kpi.value1, kpi.value2, kpi.name)
print(f"calc_accel_5: {round(time.time()- s, 2)} sec.")

运行结果

calc in fnc : 3.07 sec.
calc_raw_1  : [原代码未运行,此处省略]
calc_accel_1: 2.56 sec.
calc_accel_2: 2.59 sec.
calc_accel_3: 2.62 sec.
calc_accel_4: 5.31 sec.
calc_accel_5: 3.19 sec.

可行优化方案

1. 修正Numba jitclass的类型匹配问题

当前Kpi_jit类的类型定义和实际传入参数不匹配:value1定义为float32,但传入的x1=5.2是64位浮点数;value2定义为int32,传入的x2=989是Python int,类型转换会带来额外开销。修正类型匹配后,jitclass的性能会大幅提升:

spec = [
    ('name', types.unicode_type),
    ('value1', types.float64),  # 和x1的float64类型匹配
    ('value2', types.int64)     # 和x2的int64类型匹配
]

@jitclass(spec)
class Kpi_jit:
    def __init__(self, name, value1, value2):
        self.name = name
        self.value1 = value1
        self.value2 = value2

修正后,calc_accel_4的运行速度会接近普通Numba加速函数的水平。

2. 提前解构dataclass变量,避免循环内属性访问

如果不想改动dataclass定义,可以在循环前提前解构变量,缓存起来避免每次循环都访问dataclass属性:

# 只解构一次,放在循环外面
kpi_v1, kpi_v2, kpi_name = kpi.value1, kpi.value2, kpi.name

s = time.time()
for i in range(10000000):
    temp1, temp2 = calc_accel_5(kpi_v1, kpi_v2, kpi_name)
print(f"calc_accel_5_opt: {round(time.time()- s, 2)} sec.")

这样能消除循环内dataclass属性访问的开销,性能会和calc_accel_1持平。

3. 使用Numba支持的类型化字典

Numba在nopython模式下支持固定类型的字典,只要提前定义好键值类型,就能兼顾变量分组管理和性能:

from numba import typed

# 创建类型匹配的字典
kpi_dict = typed.Dict.empty(
    key_type=types.unicode_type,
    value_type=types.union_type(types.unicode_type, types.float64, types.int64)
)
kpi_dict['name'] = filter_text
kpi_dict['value1'] = x1
kpi_dict['value2'] = x2

@njit
def calc_accel_dict(kpi_dict):
    if kpi_dict['name'] == 'AS':
        return (kpi_dict['value1'] * kpi_dict['value2']) * 2, (kpi_dict['value1'] + kpi_dict['value2']) * 2
    else:
        return (kpi_dict['value1'] * kpi_dict['value2']) * 3, (kpi_dict['value1'] + kpi_dict['value2']) * 4

注意:尽量保持字典值的类型一致,这样性能会更好;如果必须存多种类型,用union_type声明。

4. 批量向量化计算(大数据场景首选)

如果实际处理的是批量数据,直接用NumPy数组做向量化计算,性能远超单循环:

# 生成批量数据
x1_arr = np.full(10000000, 5.2, dtype=np.float64)
x2_arr = np.full(10000000, 989, dtype=np.int64)
filter_arr = np.full(10000000, 'AB', dtype='U2')

@njit
def calc_vectorized(x1, x2, stext):
    temp1 = np.empty_like(x1, dtype=np.float64)
    temp2 = np.empty_like(x1, dtype=np.float64)
    for i in range(x1.shape[0]):
        if stext[i] == 'AS':
            temp1[i] = (x1[i] * x2[i]) * 2
            temp2[i] = (x1[i] + x2[i]) * 2
        else:
            temp1[i] = (x1[i] * x2[i]) * 3
            temp2[i] = (x1[i] + x2[i]) * 4
    return temp1, temp2

s = time.time()
temp1_arr, temp2_arr = calc_vectorized(x1_arr, x2_arr, filter_arr)
print(f"calc_vectorized: {round(time.time()- s, 2)} sec.")

向量化计算能充分利用CPU缓存和Numba的编译优化,性能提升非常显著。


内容的提问来源于stack exchange,提问作者HakanA

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 14:59:51