从MySQL数据库实时更新Seaborn热力图失效问题
问题描述
- 已完成部分:将8*8的numpy数组上传至MySQL数据库,数据库表设计为含64个float字段、id为主键。
- 核心问题:通过Seaborn和Matplotlib从数据库读取数据并实时更新热力图。测试numpy随机值生成时实时更新正常,但使用自定义数据库检索代码时,即使数据库持续插入新值,热力图仍保持旧值不变。
数据上传文件代码
import mysql.connector import serial import numpy as np from matplotlib import pyplot as plt import seaborn as sns import pandas as pd import openpyxl from multiprocessing import Process, cpu_count, Pool from matplotlib.animation import FuncAnimation ser = serial.Serial('', ) ser.close() print(ser.name) temarray = [] host_str = "" user_str = "" password_str = "" dbname = "" pydb = mysql.connector.connect(host=host_str, user=user_str, password=password_str, database=dbname) sql_insert_stmt = "insert into sensor_reads(value0, value1, value2, value3, value4, value5, value6, value7, value8, value9, value10, value11, value12, value13, value14, value15, value16, value17, value18, value19, value20, value21, value22, value23, value24, value25, value26, value27, value28, value29, value30, value31, value32, value33, value34, value35, value36, value37, value38, value39, value40, value41, value42, value43, value44, value45, value46, value47, value48, value49, value50, value51, value52, value53, value54, value55, value56, value57, value58, value59, value60, value61, value62, value63) values (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)" sql_retrieve_stmt = "select value0, value1, value2, value3, value4, value5, value6, value7, value8, value9, value10, value11, value12, value13, value14, value15, value16, value17, value18, value19, value20, value21, value22, value23, value24, value25, value26, value27, value28, value29, value30, value31, value32, value33, value34, value35, value36, value37, value38, value39, value40, value41, value42, value43, value44, value45, value46, value47, value48, value49, value50, value51, value52, value53, value54, value55, value56, value57, value58, value59, value60, value61, value62, value63 from sensor_reads ORDER BY id DESC LIMIT 0, 1;" cursor1 = pydb.cursor() cursor2 = pydb.cursor() def animate(list_corr0): ax = sns.heatmap(list_corr0, annot=True, fmt='.1f', vmin=0, vmax=300, linewidth=0.5) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') def readcom(): with serial.Serial(' ', ) as ser: while True: line0 = ser.readline() line0 = ser.readline() if line0 != None: line = line0[0: -3] print(line) line = line.decode('utf-8') print(line) line = line.split(",") print(line) list = np.array(line) print(list) list = list.astype(np.float64) list = list / 10 print(list) print("Size: ", list.shape[0]) data = (list[0], list[1], list[2], list[3], list[4], list[5], list[6], list[7], list[8], list[9], list[10], list[11], list[12], list[13], list[14], list[15], list[16], list[17], list[18], list[19], list[20], list[21], list[22], list[23], list[24], list[25], list[26], list[27], list[28], list[29], list[30], list[31], list[32], list[33], list[34], list[35], list[36], list[37], list[38], list[39], list[40], list[41], list[42], list[43], list[44], list[45], list[46], list[47], list[48], list[49], list[50], list[51], list[52], list[53], list[54], list[55], list[56], list[57], list[58], list[59], list[60], list[61], list[62], list[63] ) cursor1.execute(sql_insert_stmt, data) pydb.commit() def main(): process1 = Process(target=readcom) process1.start() process1.join() if __name__ == '__main__': main() ser.close() pydb.close()
检索和绘图文件代码
import mysql.connector import numpy as np from matplotlib import pyplot as plt import seaborn as sns import pandas as pd import openpyxl from multiprocessing import Process, cpu_count, Pool import matplotlib.animation as animation import multiprocessing temarray = [] host_str = "" user_str = "" password_str = "" dbname = "" pydb = mysql.connector.connect( host=host_str, user=user_str, password=password_str, database=dbname) sql_insert_stmt = "insert into sensor_reads(value0, value1, value2, value3, value4, value5, value6, value7, value8, value9, value10, value11, value12, value13, value14, value15, value16, value17, value18, value19, value20, value21, value22, value23, value24, value25, value26, value27, value28, value29, value30, value31, value32, value33, value34, value35, value36, value37, value38, value39, value40, value41, value42, value43, value44, value45, value46, value47, value48, value49, value50, value51, value52, value53, value54, value55, value56, value57, value58, value59, value60, value61, value62, value63) values (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)" sql_retrieve_stmt = "select value0, value1, value2, value3, value4, value5, value6, value7, value8, value9, value10, value11, value12, value13, value14, value15, value16, value17, value18, value19, value20, value21, value22, value23, value24, value25, value26, value27, value28, value29, value30, value31, value32, value33, value34, value35, value36, value37, value38, value39, value40, value41, value42, value43, value44, value45, value46, value47, value48, value49, value50, value51, value52, value53, value54, value55, value56, value57, value58, value59, value60, value61, value62, value63 from sensor_reads ORDER BY id DESC LIMIT 0, 1;" cursor1 = pydb.cursor() def retrieve(): # listfromdb = np.zeros(64) cursor2 = pydb.cursor() cursor2.execute(sql_retrieve_stmt) result = cursor2.fetchall() result = np.array(result) temparray1 = result.reshape(8, 8) temparray2 = np.array(temparray1) temparray3 = temparray2.astype(np.float32) for i in range(temparray3.shape[0]): for j in range(temparray3.shape[1]): temparray3[i, j] = temparray3[i, j] cursor2.execute(sql_retrieve_stmt) result = cursor2.fetchall() result = np.array(result) temparray1 = result.reshape(8, 8) temparray2 = np.array(temparray1) temparray3 = temparray2.astype(np.float32) for i in range(temparray3.shape[0]): for j in range(temparray3.shape[1]): temparray3[i, j] = temparray3[i, j] listfromdb = temparray3.astype(float) cursor2.close() return listfromdb def animate_heat_map(): fig = plt.figure() nx = ny = 8 data = retrieve() ax = sns.heatmap(data, annot=True, vmin = 0, vmax=300) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') def init(): plt.clf() ax = sns.heatmap(data, annot=True, vmin = 0, vmax=300) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') def animate(i): plt.clf() data = retrieve() ax = sns.heatmap(data, annot=True, vmin = 0, vmax=300) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') anim = animation.FuncAnimation(fig, animate, init_func=init, interval=1000) plt.show() def main(): process1 = Process(target=retrieve) process2 = Process(target=animate_heat_map) process1.start() process2.start() process1.join() process2.join() if __name__ == '__main__': main()
问题原因及修复方案
核心问题点
- 数据库连接复用异常:全局创建的数据库连接在子进程中复用,可能导致连接缓存或状态异常,无法获取最新数据。
- 冗余代码干扰:
retrieve函数内重复执行两次查询,且存在无意义的循环赋值,可能导致数据更新逻辑混乱。 - 多进程调用冗余:
main函数中启动的retrieve进程仅执行一次查询,对实时更新无帮助,反而占用资源。
修复代码
1. 重构retrieve函数
每次查询都创建新的数据库连接,避免跨进程复用问题,同时删除冗余代码:
def retrieve(): # 每次查询创建独立连接 pydb = mysql.connector.connect( host=host_str, user=user_str, password=password_str, database=dbname) cursor = pydb.cursor() cursor.execute(sql_retrieve_stmt) result = cursor.fetchall() # 关闭连接释放资源 cursor.close() pydb.close() # 简化数据转换逻辑 result = np.array(result) return result.reshape(8, 8).astype(float)
2. 优化动画绘制逻辑
避免每次刷新都重新创建轴对象,直接更新数据提升性能:
def animate_heat_map(): fig = plt.figure() # 初始化热力图 data = retrieve() ax = sns.heatmap(data, annot=True, vmin=0, vmax=300) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') def animate(i): nonlocal ax # 清除当前轴内容 ax.clear() # 获取最新数据 data = retrieve() # 重新绘制热力图 ax = sns.heatmap(data, annot=True, vmin=0, vmax=300, ax=ax) ax.invert_yaxis() ax.set(xlabel='Column number', ylabel='Row number') # 移除init_func,直接用初始绘制的内容 anim = animation.FuncAnimation(fig, animate, interval=1000) plt.show()
3. 简化main函数
无需启动冗余进程,直接调用动画函数即可:
def main(): animate_heat_map()
内容的提问来源于stack exchange,提问作者RedRabbit
相关产品推荐
相关产品推荐

