Zyphra/zuna

GitHub: Zyphra/zuna

ZUNA1.1 是一个开源的 EEG 基础模型,能够基于 3D 头皮坐标对脑电信号进行去噪、缺失通道重构和稀疏导联上采样。

Stars: 316 | Forks: 57

ZUNA1.1 — Thought to Text

# ZUNA1.1:一个灵活的 EEG 基础模型 [![HuggingFace ZUNA](https://img.shields.io/badge/HuggingFace-ZUNA1.1-FFD21E?logo=huggingface&logoColor=black&labelColor=555555)](https://huggingface.co/Zyphra/ZUNA1.1) [![PyPI](https://img.shields.io/pypi/v/zuna?label=pypi&logo=pypi&logoColor=white)](https://pypi.org/project/zuna/) [![加入我们的 Discord](https://img.shields.io/discord/1304567558682443806?label=Join%20our%20Discord&logo=discord&logoColor=black)](https://discord.gg/ZF7BCgjAcC) [![许可证](https://img.shields.io/badge/License-Apache%202.0-blue.svg)](LICENSE) **ZUNA1.1** 是 Zyphra 为 EEG 开发的开源基础模型。它能够重构含噪或缺失的通道、对记录进行去噪,并将稀疏的电极分布上采样为更密集的分布。由于它基于每个电极的 **3D 头皮坐标** 而非固定的通道列表进行条件设定,因此它无需重新训练即可应用于几乎任何导联组合——从 4 通道的 Muse 头带到 256 通道的研究级电极帽,甚至能够生成从未记录过的电极位置的信号。 - **去噪** 现有 EEG 通道 - **重构** 缺失或脱落的通道 - **上采样** 稀疏导联组合——根据头皮坐标预测新通道 ZUNA1.1 是一个拥有 3.8 亿参数的 diffusion autoencoder,在约 350 万小时的公开 EEG 数据上训练而成。它仅有 3.8 亿参数,只需 **<1 GB 的 VRAM**,在消费级 GPU 或 Mac(Apple Silicon)上运行速度极快,并在 CPU 上也能获得可接受的效果。 ## ☁️ 在浏览器中试用 — Zyphra Cloud 无需安装,无需 GPU,无需编写代码。上传 EEG 记录文件(`.fif`)——或使用提供的示例——进入 **[Zyphra Cloud EEG Playground](https://cloud.zyphra.com)**,标记噪声片段(手动或自动选择),即可直接在浏览器中进行去噪或上采样。我们在我们的服务器上托管模型并运行推理;您的会话结束后不会保留任何内容,我们也不会使用用户数据进行训练。 ## ZUNA1.1 的新特性 ZUNA1.1 保留了原始 [ZUNA1](https://huggingface.co/Zyphra/ZUNA) 的架构,但在训练上针对现实世界数据实现了更高的灵活性和鲁棒性,同时匹配甚至超越了 ZUNA1 的重构质量: 1. **可变长度输入(0.5–30 秒):** 根据每个训练样本对片段长度进行采样(对齐到 0.125 秒的 token 网格),而不是仅仅使用固定的 5 秒窗口,因此同一个模型无需重新配置即可服务于 0.5 秒的试验片段或 30 秒的连续片段。 2. **更丰富的重构任务组合:** 在 **四种** 真实的通道丢失模式上进行训练(参见 [训练](#training)),而不是单一的随机丢弃方案,涵盖了现实世界中 EEG 实际受损的多种方式。 3. **质量感知预处理和更大规模的数据集:** 通过逐通道、逐秒的质量评分,可以从部分含噪通道中恢复信号(旧的全记录 pipeline 会直接丢弃这些信号),从而将数据集从约 200 万小时增加到约 350 万小时。每条记录提供两种滤波变体(0.1–45 Hz 带通和更轻量的 0.01 Hz 高通 + notch),使模型能够泛化到异构的预处理方式。 ## 架构

ZUNA1.1 architecture
ZUNA1.1 is a transformer encoder–decoder diffusion autoencoder trained to reconstruct masked EEG channels. The main changes from ZUNA1 improve training stability (e.g. additional normalization layers).

ZUNA 将每个 EEG 通道切分为短小的 **0.125 秒片段(在 256 Hz 下为 32 个采样点)**,将其转换为连续值的 token,并按通道 × 时间的顺序进行序列化。其核心思想在于位置编码:每个 token 都带有一个 **基于 (x, y, z, t) 的 4D rotary positional encoding** ——即电极的 3D 头皮坐标加上其粗略的时间索引。因为告诉模型通道所在位置的是 *位置* 而不是数组索引,所以 ZUNA 是 **与通道无关的 (channel-agnostic)**:它接受任何布局中任意数量的电极,并且能够在从未记录过的位置合成信号(基于位置的任意上采样)。编码器将信号压缩为潜变量,通过 adaptive-RMS norm 对解码器进行条件设定;解码器使用 rectified-flow 目标进行训练。 ## 训练 ZUNA1.1 在由 **四种通道丢失方案** 组成的混合数据上进行了训练,每种方案都捕捉了 EEG 受损或数据缺失的不同方式: - **Whole-channel:** 整个通道被移除(稀疏导联组合、坏电极)。 - **Full-time:** 跨 *所有* 通道移除短时间片段(全信号丢失、头部运动爆发)。 - **Channel-time:** 仅从 *部分* 通道中移除相同的时间片段(在空间和时间上聚集的间隙,例如附近电极上的运动伪影)。 - **Random-uniform:** 散布在单个采样点上的缺失值(短暂的、局部的噪声,如肌肉抽搐)。

ZUNA1.1 dropout schemes
ZUNA1.1's dropout schemes are far more diverse than ZUNA1, which dropped entire channels over all time. Training across this mixture lets ZUNA1.1 handle almost arbitrary reconstructions across space and time.

## 性能 增加这种灵活性在重构质量上并没有造成明显的损失。在留出数据集的评估中,ZUNA1.1 达到了与 ZUNA1 相同甚至更好的 NMSE,并且两者都明显优于经典的球面样条插值——随着丢失通道数量的增加,这种差距不断扩大,而样条插值(仅假设空间平滑性)在这种情况下则会失效。

Reconstruction accuracy as channels drop
Reconstruction NMSE vs channel-dropout rate across four datasets — ZUNA1.1 vs ZUNA1 vs MNE spherical-spline interpolation. Lower is better. (Evaluation restricted to 5 s samples for comparison with ZUNA1.)

我们还评估了一个更贴近真实实验的设置:删除某个大脑区域的所有电极,并根据其余七个区域进行重构。ZUNA1.1 在各个区域均处于领先地位。

Reconstruction accuracy by brain region (topographic)
Per-region reconstruction NMSE (topographic view): ZUNA1.1, ZUNA1 (Δ vs ZUNA1.1), and spherical-spline. Lower/greener is better.

Reconstruction errors by brain region (bar chart)
The same per-region errors as a bar chart, averaged across four datasets; error bars show propagated standard deviation. Lower is better.

## 安装 ``` # (1) 克隆 repo(用于教程 + 示例数据) git clone https://github.com/Zyphra/zuna.git && cd zuna # (2) 安装 zuna pip install zuna ``` 或者以开发模式安装: ``` git clone https://github.com/Zyphra/zuna.git && cd zuna pip install -e . ``` ### GPU 支持 (PyTorch + CUDA) `zuna` 通过 PyTorch 在 GPU 上运行,并且 **PyPI 无法为您选出与您的 GPU 驱动相匹配的 PyTorch 构建版本**。如果自动安装的 `torch` 是为比您的 NVIDIA 驱动支持的版本更新的 CUDA 版本构建的,PyTorch 会静默回退到 CPU(非常慢),并发出诸如 `No CUDA runtime is found` / `CUDA initialization: The NVIDIA driver on your system is too old` 的警告。 为了避免这种情况,请在安装 `zuna` **之前** 安装与您的驱动相匹配的 `torch` 构建版本。检查您的驱动支持的 CUDA 版本(`nvidia-smi` 的右上角),然后安装对应的 wheel——例如,对于 CUDA 12.8: ``` # 1. 安装与您的驱动 CUDA 版本匹配的 torch build(参见 `nvidia-smi`)。 # 以 CUDA 12.8 为例 — 使用 cu121 / cu124 / cu126 / cu128 来匹配您的版本: pip install torch --index-url https://download.pytorch.org/whl/cu128 # 2. 然后安装 zuna(它将使用您已安装的 torch) pip install zuna ``` 如果您已经安装了 `zuna` 并且它正在 CPU 上运行,请通过重新安装匹配的 torch 来修复此问题: ``` pip install --force-reinstall torch --index-url https://download.pytorch.org/whl/cu128 ``` 验证 GPU 访问权限: ``` python -c "import torch; print(torch.__version__, torch.version.cuda, torch.cuda.is_available())" # 预期得到 CUDA build(例如 ...+cu128)和 `True`。 ``` ## 快速开始 `tutorials/run_zuna_pipeline.py` 是一个完整且可编辑的示例。它从输入目录读取 `.fif` 文件,使用模型重构选定的单元,并将 `.fif` 文件写回(无需进行 `.pt` 往返转换)。编辑顶部的常量,然后运行: ``` python tutorials/run_zuna_pipeline.py ``` 首次运行时会自动从 HuggingFace 下载模型权重。输出结果将保存在 `OUTPUT_DIR` 下: ``` 2_fif_output/ full_reconstruction/_raw.fif # model output everywhere hybrid/_raw.fif # original input, model output ONLY on the inferred cells hybrid/_mask.npz # per (channel, token) mask of what was inferred figures/ __full_reconstruction.png # full-duration input-vs-reconstruction overlay __hybrid.png ``` ## 重构 `.fif` 文件:`reconstruct_fif` 运行器调用了一个函数,您也可以直接使用该函数: ``` from zuna import reconstruct_fif reconstruct_fif( input_dir="path/to/fif/input", output_dir="path/to/fif/output", figures_dir="path/to/figures", gpu_device=0, # GPU id, or "" for CPU highpass_hz=0.5, # highpass applied before the model (None to skip) montage="standard_1020", # used only to add positions when a .fif lacks them ) ``` **重构 mask 是以下所有来源的并集 (UNION)** ——任何来源都不会覆盖其他来源,因此您可以自由地将自动检测与手动选择结合起来。 ### 自动:MNE 坏通道 + `BAD_` 标注 如果您的 `.fif` 已经使用 MNE 标记了坏数据,ZUNA 无需额外参数即可使用它: - **`info['bads']`** —— 任何被标记为坏的通道都会被完全重构(始终使用)。 - **`BAD_*` 标注** —— 被标注为坏的时间跨度(跨越所有通道)将被重构。可通过 `use_fif_annotations`(默认为 `True`)进行切换。 ``` reconstruct_fif(..., use_fif_annotations=True) # import the .fif's own BAD_ time annotations ``` ### 修复特定通道(即使未标记为坏) 指定要完全重构的通道名称,无论它们是否在文件中被标记为坏: ``` reconstruct_fif(..., repair_channels=["Cz", "Fz"]) ``` ### 添加通道 / 上采样导联组合 在头皮位置预测全新的通道。传入 **names** 以精确添加这些通道,或者传入一个 **integer** 以自动添加通道,使通道总数达到该数值,并将它们放置在远离现有电极的地方: ``` reconstruct_fif(..., target_channel_count=["Fz", "Pz"]) # add these exact channels reconstruct_fif(..., target_channel_count=40) # auto-upsample to 40 channels total ``` ### 重构手动选择的时间片段 传入一个元组列表,单位为 **相对于数据的秒数**。一个包含 2 个元素的元组表示该时间段在 **所有** 通道上都被标记为坏;一个包含 3 个元素的元组则将其限制在 **某一个** 通道: ``` reconstruct_fif(..., bad_segments=[ (5, 6), # 5–6 s bad on ALL channels (10, 11, "C3"), # 10–11 s bad on C3 only (10, 11, "C4"), # 10–11 s bad on C4 only ]) ``` ### 从 UI / 外部 mask 驱动 对于 UI(或任何外部工具),提供一个包含每个文件 mask 的目录。每个 `_mask.npz` 包含一个 `(channel × token)` 的布尔数组(每 `num_fine_time_pts` = 32 个采样点 ≈ 0.125 秒对应一列;也接受采样分辨率),以及 `ch_names` 和 `sfreq`。它会与上述所有内容取并集: ``` reconstruct_fif(..., mask_dir="path/to/masks") ``` 您可以使用辅助工具 `zuna.write_bad_mask(...)` 根据坏通道 + 时间片段构建这样的 mask,它会生成与重构器写入 `hybrid/_mask.npz` 相同的格式——因此 UI 可以对其进行完整的往返处理。 ## 设置导联组合 ZUNA 需要 3D 电极位置。如果您的 `.fif` 未携带导联组合: ``` import mne raw = mne.io.read_raw_fif("data.fif", preload=True) raw.set_montage(mne.channels.make_standard_montage("standard_1005")) raw.save("data_with_montage.fif", overwrite=True) ``` 任何具有已知位置的导联组合均可工作——从消费级头戴设备的布局到标准的 8/16/32/64 通道研究级导联组合,甚至高达 256 通道的系统。 ## 引用 技术白皮书即将发布。如果您在研究中发现 ZUNA 有用,请按此进行引用。 如果组织或研究人员有兴趣与 Zyphra 合作,针对特定需求或用例改进未来版本,请联系 bci@zyphra.com。 ## 免责声明 本软件及相关服务(统称“服务”)仅供研究使用,不用于任何疾病或健康状况的诊断、治愈、缓解、治疗或预防。这些服务尚未在任何医疗或临床用途中得到验证。通过服务提供的信息仅供参考,不能替代任何专业的医疗或保健建议。我们不保证通过服务提供的任何信息对您而言是准确、完整或有用的。您对此类信息的任何依赖均需严格自行承担风险。
标签:EEG(脑电图), Python, 人工智能, 信号去噪, 凭据扫描, 基础模型, 数据重建, 无后门, 用户模式Hook绕过, 逆向工具