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

Python+Matplotlib病毒传播模拟如何新增实时统计曲线独立窗口

实现方案

实现思路

  • 新增第二个独立的Matplotlib窗口作为统计面板,初始化4条和人群类型颜色对应的曲线(健康绿、感染红、治愈蓝、死亡黑)
  • 新增全局数组存储时间序列和对应时刻的四类人群数量,用于更新曲线数据源
  • 把曲线更新逻辑嵌入原有的动画回调函数中,保证粒子运动和统计曲线完全同步
  • 保留两个动画的对象引用,避免被Python垃圾回收导致动画停止

完整修改后代码

"""
Description: 模拟病毒在群体中的传播,包含粒子运动窗口和人群数量实时统计曲线窗口
"""

# -------------------- IMPORTS --------------------
from random import randint, random
from matplotlib import pyplot as plt, animation as anim
import math


# --------------------  GLOBAL VARIABLES --------------------
number_of_dots = 100  # 生成的粒子总数
shape = "o"  # 粒子样式,最小距离<3时建议用'.'

HEIGHT_WIDTH = 100  # 粒子窗口宽高(必须为正方形)
BORDER_MIN = 1  # 粒子距边框最小距离
BORDER_MAX = HEIGHT_WIDTH - 1  # 粒子距边框最大距离

minimal_distance = 3  # 初始化和感染判定的最小距离
time = 0  # 初始时间
time_step = 0.1  # 模拟时间步长

transmission_rate = 0.7  # 接触后感染概率
time_to_cure = 40  # 感染后自愈所需时间
time_before_being_contagious_again = 40  # 治愈后可再次感染的间隔时间
virus_mortality = 0.0005  # 感染后每帧死亡概率

# 统计曲线相关全局变量
time_history = []
healthy_history = []
infected_history = []
cured_history = []
dead_history = []
figureStats = None
ax_stats = None
healthy_line = None
infected_line = None
cured_line = None
dead_line = None


# -------------------- CLASSES & METHODS --------------------
class Dot:
    def __init__(self, x: int, y: int):
        """Dot类构造函数
        Args:
            x (int): 粒子横坐标
            y (int): 粒子纵坐标
        """
        self.x = x
        self.y = y
        self.velx = (random() - 0.5) / 5
        self.vely = (random() - 0.5) / 5
        self.is_infected = False
        self.infected_at = -1
        self.has_been_infected = False
        self.cured_at = -1

    def init_checker(x: int, y: int, already_used_coords: list):
        """检查初始化粒子是否与已有粒子距离过近
        Args:
            x (int): 待初始化粒子横坐标
            y (int): 待初始化粒子纵坐标
            already_used_coords (list): 已占用坐标列表
        Returns:
            boolean: 该位置是否可初始化粒子
        """
        for coord in already_used_coords:
            if Dot.get_distance(coord[0], x, coord[1], y) < minimal_distance:
                return False
        return True

    def get_distance(x1: float, y1: float, x2: float, y2: float):
        """计算两个粒子的距离
        Args:
            x1 (float): 第一个粒子横坐标
            y1 (float): 第一个粒子纵坐标
            x2 (float): 第二个粒子横坐标
            y2 (float): 第二个粒子纵坐标
        Returns:
            float: 两个粒子的欧氏距离
        """
        return math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)

    def initalize_multiple_dots():
        """批量初始化粒子
        Returns:
            list: 初始化完成的粒子列表
        """
        dots = []
        already_used_coords = []

        while len(dots) != number_of_dots:
            randx = randint(BORDER_MIN, BORDER_MAX)
            randy = randint(BORDER_MIN, BORDER_MAX)
            if Dot.init_checker(randx, randy, already_used_coords):
                dot = Dot(randx, randy)
                already_used_coords.append((randx, randy))
            else:
                continue
            dots.append(dot)
        print(f"区域内共有{len(dots)}个粒子")
        return dots

    def move(self):
        """粒子移动逻辑,包含边界碰撞、方向随机变化、感染判定、死亡判定"""
        global dots, dead_dots

        if random() < 0.96:
            self.x = self.x + self.velx
            self.y = self.y + self.vely
        else:
            self.x = self.x + self.velx
            self.y = self.y + self.vely
            self.velx = (random() - 0.5) / (2 / (time_step + 1))
            self.vely = (random() - 0.5) / (2 / (time_step + 1))

        # 边界碰撞反弹
        if self.x >= BORDER_MAX:
            self.x = BORDER_MAX
            self.velx = -1 * self.velx
        if self.x <= BORDER_MIN:
            self.x = BORDER_MIN
            self.velx = -1 * self.velx
        if self.y >= BORDER_MAX:
            self.y = BORDER_MAX
            self.vely = -1 * self.vely
        if self.y <= BORDER_MIN:
            self.y = BORDER_MIN
            self.vely = -1 * self.vely

        # 感染判定
        if (
            random() < transmission_rate
            and not self.has_been_infected
            and not self.is_infected
        ):
            for dot in dots:
                if (
                    dot.is_infected
                    and Dot.get_distance(self.x, self.y, dot.x, dot.y)
                    < minimal_distance
                ):
                    self.is_infected = True
                    self.infected_at = time
                    break
        
        # 死亡判定
        if self.is_infected and random() < virus_mortality:
            dead_dots.append(self)
            dots.remove(self)

    def move_all(dots: list):
        """批量更新粒子状态,同时更新粒子窗口和统计曲线
        Args:
            dots (list): 存活粒子列表
        """
        global healthy_dots, infected_dots, cured_dots, time
        global figureStats, ax_stats, healthy_line, infected_line, cured_line, dead_line
        global time_history, healthy_history, infected_history, cured_history, dead_history

        for dot in dots:
            dot.move()
            # 自愈判定
            if (
                dot.is_infected
                and dot.infected_at != -1
                and dot.infected_at + time_to_cure < time
            ):
                dot.is_infected = False
                dot.has_been_infected = True
                dot.cured_at = time
            # 免疫失效判定
            if (
                dot.has_been_infected
                and dot.cured_at != -1
                and dot.cured_at + time_before_being_contagious_again < time
            ):
                dot.has_been_infected = False
                dot.infected_at = -1
                dot.cured_at = -1

        # 更新粒子窗口显示
        cnt_healthy = len([dot for dot in dots if not dot.is_infected and not dot.has_been_infected])
        cnt_infected = len([dot for dot in dots if dot.is_infected])
        cnt_cured = len([dot for dot in dots if dot.has_been_infected])
        cnt_dead = len(dead_dots)

        healthy_dots.set_data(
            [dot.x for dot in dots if not dot.is_infected and not dot.has_been_infected],
            [dot.y for dot in dots if not dot.is_infected and not dot.has_been_infected],
        )
        infected_dots.set_data(
            [dot.x for dot in dots if dot.is_infected],
            [dot.y for dot in dots if dot.is_infected],
        )
        cured_dots.set_data(
            [dot.x for dot in dots if dot.has_been_infected],
            [dot.y for dot in dots if dot.has_been_infected],
        )

        plt.title(
            f"健康: {cnt_healthy}   |   感染: {cnt_infected}   |   治愈: {cnt_cured}   |   死亡: {cnt_dead}",
            color="black",
        )

        # 更新统计曲线
        time_history.append(time)
        healthy_history.append(cnt_healthy)
        infected_history.append(cnt_infected)
        cured_history.append(cnt_cured)
        dead_history.append(cnt_dead)

        healthy_line.set_data(time_history, healthy_history)
        infected_line.set_data(time_history, infected_history)
        cured_line.set_data(time_history, cured_history)
        dead_line.set_data(time_history, dead_history)

        # 动态扩展X轴范围
        ax_stats.set_xlim(0, max(time_history) + 1)
        figureStats.canvas.draw()

        time += time_step


# -------------------- MAIN FUNCTION --------------------
def main():
    global dots, dead_dots, axes
    global figureStats, ax_stats, healthy_line, infected_line, cured_line, dead_line

    # 初始化粒子
    dots = Dot.initalize_multiple_dots()
    random_infected = randint(0, len(dots) - 1)
    dots[random_infected].is_infected = True
    dots[random_infected].infected_at = time
    dead_dots = []

    # 创建粒子模拟窗口
    figureDots = plt.figure(facecolor="white", figsize=(5,5))
    axes = plt.axes(xlim=(0, HEIGHT_WIDTH), ylim=(0, HEIGHT_WIDTH))
    plt.axis('off')

    global healthy_dots, infected_dots, cured_dots
    healthy_dots = axes.plot([], [], f"g{shape}")[0]
    infected_dots = axes.plot([], [], f"r{shape}")[0]
    cured_dots = axes.plot([], [], f"b{shape}")[0]

    # 创建统计曲线窗口
    figureStats = plt.figure(facecolor="white", figsize=(7,4))
    ax_stats = plt.axes()
    ax_stats.set_xlabel("模拟时间")
    ax_stats.set_ylabel("人数")
    ax_stats.set_ylim(0, number_of_dots + 10)
    healthy_line, = ax_stats.plot([], [], 'g-', label='健康')
    infected_line, = ax_stats.plot([], [], 'r-', label='感染')
    cured_line, = ax_stats.plot([], [], 'b-', label='治愈')
    dead_line, = ax_stats.plot([], [], 'k-', label='死亡')
    ax_stats.legend()
    ax_stats.grid(alpha=0.3)

    # 绑定两个窗口的动画,必须保留两个动画的引用否则会被回收
    anim_dots = anim.FuncAnimation(figureDots, lambda z: Dot.move_all(dots), frames=60, interval=5)
    anim_stats = anim.FuncAnimation(figureStats, lambda z: None, frames=60, interval=5)

    plt.show()


# -------------------- MAIN CALL --------------------
if __name__ == "__main__":
    main()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 01:48:03