Dokümantasyon
Kur, veriyi doğrula, eğit, değerlendir
Bu sayfadaki her komut güncel depoya karşı çalışır. Yer tutucu çıktı yok, uydurma bayrak yok — bir şey henüz çalışmıyorsa bu açıkça belirtilir.
Gereksinimler
Python 3.9 veya üzeri, PyTorch 2.0 veya üzeri. Yığının geri kalanı requirements.txt dosyasından gelir: BART decoder için transformers, EEG ön işleme için mne, ayrıca değerlendirme için numpy, scipy, pandas, lightning, h5py, jiwer ve python-Levenshtein.
--quick-test ve CTC modeli için bir CPU veya MPS destekli Apple Silicon Mac yeterlidir. BART modeli --fp16 ile eğitilir, bu da CUDA destekli bir NVIDIA GPU gerektirir; bu GPU süresi henüz sağlanamadığı için NEST v2 BART tam olarak eğitilmemiştir.
Kurulum
Depoyu klonlayın, bir sanal ortam oluşturun ve bağımlılıkları kurun.
git clone https://github.com/wazder/nest cd nest python -m venv .venv source .venv/bin/activate pip install -r requirements.txt
PyPI paketi yoktur. pip install nest diye bir şey yoktur ve pip install -e . henüz çalışmaz: pyproject.toml dosyasında bir [project] tablosu yoktur, bu yüzden pip'in derleyeceği bir şey yoktur. Betikleri, bu sayfada gösterildiği gibi, klonlanmış bir depodan doğrudan çalıştırın.
Veri
NEST, ZuCo korpusu üzerinde eğitilir. osf.io/q3zws adresinden erişim talep edin ve çıkarılan pickle dosyalarını, yükleyicinin beklediği düzenle eşleşecek şekilde yerleştirin, her görev için bir pickle dosyası:
ZuCo_Dataset/ZuCo/task1-SR/pickle/task1-SR-dataset.pickle ZuCo_Dataset/ZuCo/task2-NR/pickle/task2-NR-dataset.pickle ZuCo_Dataset/ZuCo/task3-TSR/pickle/task3-TSR-dataset.pickle
Veri setinin yüklendiğini ve beklenen cümle-denek çifti sayısını bildirdiğini doğrulayın:
python src/data/zuco_pickle_dataset.py ZuCo_Dataset/ZuCo
Hızlı test
Tam bir koşudan önce, veri hattının uçtan uca çalıştığını doğrulayın. --quick-test, 50 örnek üzerinde 2 epoch eğitir ve görülmemiş test değerlendirmesini atlar.
python scripts/train_nest_v2.py --quick-test --model ctc
Eğitim
Eğitim giriş noktası scripts/train_nest_v2.py dosyasıdır. --model ile seçilen ctc ve bart olmak üzere iki model tipini destekler.
CTC temel modeli (doğrusal CTC başlığı, CPU, MPS veya GPU üzerinde çalışır):
python scripts/train_nest_v2.py --model ctc --epochs 200
BART decoder (en iyi kalite, --fp16 için NVIDIA GPU gerektirir):
python scripts/train_nest_v2.py \
--model bart \
--epochs 200 \
--batch-size 16 \
--fp16 \
--d-model 768 \
--num-layers 6 \
--tasks task1-SR task2-NR task3-TSR
train_nest_v2.py üzerinde var olan, argüman ayrıştırıcısına göre doğrulanmış bayraklar:
| Bayrak | Varsayılan | Anlam |
|---|---|---|
| --model | ctc | Model tipi: ctc veya bart |
| --data-dir | ZuCo_Dataset/ZuCo | ZuCo veri setinin yolu |
| --tasks | task1-SR task2-NR task3-TSR | Dahil edilecek ZuCo görevleri |
| --epochs | 200 | Eğitim epoch sayısı |
| --batch-size | 16 | Eğitim batch boyutu |
| --lr | 3e-4 | Öğrenme oranı |
| --d-model | 512 | Transformer gizli katman boyutu, 512 veya 768 |
| --num-layers | 6 | Transformer encoder katman sayısı |
| --nhead | 8 | Dikkat başlığı sayısı |
| --fp16 | off | Karma hassasiyet; CUDA gerektirir |
| --grad-accum | 4 | Gradyan biriktirme adımları |
| --patience | 20 | Erken durdurma sabrı, epoch cinsinden |
| --output-dir | none | Sonuçların yazılacağı yer |
| --resume | none | Devam edilecek checkpoint'in yolu |
| --num-workers | 0 | DataLoader worker süreç sayısı |
| --fixation | GD | EEG fiksasyon ölçüsü: FFD, TRT veya GD |
| --quick-test | off | 2 epoch, 50 örnek, test değerlendirmesi atlanır |
| --no-subject-split | off | Denekten bağımsız yerine rastgele 80/10/10 bölme |
Değerlendirme
Değerlendirme ayrı bir betik değildir. --resume'a bir checkpoint verin, son değerlendirme adımının onu yeniden yükleyebilmesi için --output-dir'ı o checkpoint'in dizinine yönlendirin ve eğitimi atlamak için --epochs 0 ayarlayın:
python scripts/train_nest_v2.py \
--resume results/nest_v2_bart_TIMESTAMP/best_model.pt \
--output-dir results/nest_v2_bart_TIMESTAMP \
--epochs 0
Sonuçlar, bir test_wer ve bir test_cer alanıyla results.json dosyasına yazılır: kelime hata oranı ve karakter hata oranı, ikisi de referans metne karşı Levenshtein düzenleme mesafesi olarak hesaplanır. BART yolu model.generate() ile çözümler, teacher forcing içermeyen serbest otoregresif çözümleme (mevcut uygulamada greedy). CTC yolu greedy olarak çözümler. Henüz bir checkpoint bulunmadığından bu yoldan henüz bir results.json üretilmemiştir.
Bulut
Yerel eğitim, BART koşusu için bir NVIDIA GPU gerektirir. notebooks/NEST_CloudTraining.ipynb, bu durum için hazır çalışır bir Colab notebook'udur. Bir A100 üzerinde, tam 200-epoch'luk bir BART koşusunun 8 ila 12 saat süreceği tahmin edilmektedir; bu tahmin, koşu tamamlanmadığı için uçtan uca ölçülmemiştir.
Depo yapısı
src/ önişleme, modeller (nest.py, nest_bart.py, nest_v2.py),
veri yükleyiciler, eğitim ve değerlendirme modülleri
scripts/ train_nest_v2.py, train_nest_bart.py, upload_to_huggingface.py
notebooks/ NEST_CloudTraining.ipynb (Colab)
papers/ NEST_manuscript.md (teknik rapor taslağı)
configs/ YAML model konfigürasyonları
tests/ 5 dosyada 80 test fonksiyonu
website/ bu site (Cloudflare Pages)
Bilinen sorunlar
- Test paketi (5 dosyada 80 test fonksiyonu), kod ve test kayması nedeniyle şu anda import edilirken hata veriyor. Henüz onarılmadı; nasıl yardımcı olabileceğinizi görmek için Katkı sayfasına bakın.
pyproject.tomldosyasında bir[project]tablosu yok, bu yüzden paketleme meta verisi eksik vepip install -e .çalışmıyor.- Eğitilmiş bir NEST v2 checkpoint'i yok. Tamamlanan tek koşu, NEST v2 değil küçük, geçici bir LSTM smoke-test modelini eğitti; loss düştü ama herhangi bir kelime hata oranı hesaplanmadı.
Bunları takip edin ve yenilerini github.com/wazder/nest/issues adresinde bildirin.
Değerlendirme protokolü taahhüdü
Görülmemiş denekler üzerinde serbest çözümlemeyle, bir gürültü kontrolüyle birlikte ölçülmeden bu sitede hiçbir kelime hata oranı yayınlanmaz. Taahhüdün tamamı ve arkasındaki gerekçe Araştırma sayfasındadır.