SCITE Reproduction Notes
========================
Paper: Li, Z., Li, Q., Zou, X., & Ren, J. (2021). "Causality extraction based on
self-attentive BiLSTM-CRF with transferred embeddings."
Neurocomputing, 423, 207–219.
DOI: 10.1016/J.NEUCOM.2020.08.078
arXiv: 1904.07629

==============================================================
PERFORMANCE TARGETS (from paper, Table 3 — triplet F1 on test set)
==============================================================

The paper runs each experiment 5 times and reports mean ± std.
A predicted triplet is correct only if BOTH spans exactly match a gold triplet.

  Model                         P               R               F
  ─────────────────────────────────────────────────────────────────
  SCITE (full, Flair)      0.8333±0.0042   0.8581±0.0021   0.8455±0.0028
  Flair-BiLSTM-CRF         0.8414±0.0079   0.8351±0.0141   0.8382±0.0092
  BERT-BiLSTM-CRF          0.8277±0.0058   0.8209±0.0093   0.8243±0.0049
  ELMo-BiLSTM-CRF          0.8361±0.0135   0.8399±0.0063   0.8379±0.0092
  Flair+CLSTM-BiLSTM-CRF   0.8403±0.0090   0.8284±0.0125   0.8343±0.0106
  BiLSTM-CRF               0.7837±0.0061   0.7932±0.0087   0.7884±0.0072
  CCNN-BiLSTM-CRF          0.8069±0.0199   0.7520±0.0227   0.7780±0.0075
  CLSTM-BiLSTM-CRF         0.8144±0.0284   0.7412±0.0073   0.7757±0.0107
  SCITE (general tagging)  0.7609±0.0145   0.7682±0.0136      —

Table 4 — tag-wise F1 for SCITE (Flair), BiLSTM-CRF, Flair-BiLSTM-CRF:
  SCITE:    C-P=0.90  C-R=0.86  C-F≈0.88   E-F≈0.91   Emb-R=0.22  Emb-F≈0.29

Ablation (Table 5, F1):
  SCITE (All = Flair+CCNN-BiLSTM-MHSA-CRF)   0.8455
  Flair-BiLSTM-MHSA-CRF  (no CCNN)           0.8438
  Flair-BiLSTM-CRF        (no CCNN, no MHSA)  0.8382
  BiLSTM-MHSA-CRF         (no Flair, no CCNN) 0.8137
  BiLSTM-CRF              (no Flair, no MHSA)  0.7884

==============================================================
OUR REPRODUCED RESULTS (10-fold CV, 1 run each)
==============================================================

  Run                                    tri_F1             tri_P              tri_R
  ─────────────────────────────────────────────────────────────────────────────────
  Flair-fast, 5 ep,  smha_concat=False   0.6716 ± 0.0387    0.6937 ± 0.0407    0.6533 ± 0.0554
  Flair-fast, 200 ep, smha_concat=True   0.7013 ± 0.0347    0.7388 ± 0.0457    0.6708 ± 0.0526
  BERT (bert-base), 5 ep, concat=False   0.7034 ± 0.0441    —                  —

  200-ep Flair-fast span F1:
    eval_f1_C:   0.8334 ± 0.0164   (paper SCITE ≈ 0.90)
    eval_f1_E:   0.8367 ± 0.0180   (paper SCITE ≈ 0.91)
    eval_f1_Emb: 0.2177 ± 0.1925   (paper SCITE ≈ 0.29)
  Early stopping fired at epoch 57.7 ± 18.4 on average (patience=30).

Gap from paper (200-ep run): Flair −14.4 pp triplet F1
Root cause: C/E span F1 is ~7 pp below paper (0.83 vs 0.90), which compounds
into a ~14 pp triplet gap because both spans must match simultaneously.
This points to the Flair-fast model substitution as the dominant remaining
factor — the -fast models were trained on a smaller corpus and produce
weaker contextual embeddings than the original news-forward-0.4.1.

==============================================================
DATASET
==============================================================

  HuggingFace: thagen/SCITE   config: causality-identification
  Train: 4,450 sentences      Test: 786 sentences

Paper dataset (SemEval 2010 T8 with extended annotation):
  Train: 4,450 sentences, 1,570 causal triplets
  Test:  804 sentences,   296 causal triplets

  Tag counts (paper Table 2):
    Train: B-C=1308  I-C=1421  B-E=1268  I-E=1230  B-Emb=55  I-Emb=55  → 5,337 total
    Test:  B-C=236   I-C=229   B-E=238   I-E=230   B-Emb=9   I-Emb=16  → 958 total

NOTE: our test split has 786 sentences vs paper's 804. The causalatee format
may have dropped 18 sentences (e.g., those without any causal tags or those
with annotation ambiguity). This affects comparability of test-set results.

==============================================================
ARCHITECTURE — WHAT THE PAPER SPECIFIES
==============================================================

Input per word: [word_emb (300) | char_CNN (30) | Flair (2048)]  = 2378-dim
  - Word embeddings: 300-D Komninos & Manandhar (wiki-extvec), frozen
  - Character embeddings: 30-D, uniform init [−√(3/30), +√(3/30)]
  - Character CNN: 1 layer, 30 filters, kernel size 3
  - Flair: news-forward (1024-D) + news-backward (1024-D) stacked → 2048-D

BiLSTM: 1-layer, hidden size 256 per direction → 512-D output

Multihead Self-Attention (MHSA) — CONCATENATION variant:
  - h = 3 heads, dv = 8 per head → attention output M has dim 24
  - H̃ = concat([H, M])  →  dim 512 + 24 = 536
  - Then H̃ → linear → 7 tag scores (NOT a residual add)
  - This corresponds to smha_concat=True in our SCITEConfig

CRF decoder: standard linear-chain, 7 tags (O B-C I-C B-E I-E B-Emb I-Emb)

Training:
  - Optimizer: Nadam, lr = 0.001
  - LR scheduler: halve if training loss does not fall for > 10 epochs
  - Dropout: variational dropout, rate 0.5 (on CCNN output and LSTM output)
  - Gradient clipping: threshold 5.0
  - Batch size: 16
  - Max epochs: 200 — pick checkpoint with highest 10-fold CV F1
  - Evaluation protocol: 5 independent runs, report mean ± std

==============================================================
KNOWN DEVIATIONS IN OUR REPRODUCTION
==============================================================

1. SMHA VARIANT — FIXED  [contributed ~3 pp; default now smha_concat=True]
   Paper uses the CONCATENATION variant: H̃ = [H; M] → linear → tags.
     h=3, dv=8 → concat dim 24, total 536 to CRF.
   Our default: smha_concat=False → residual addition (attention output added
     to LSTM output, not concatenated), with 4 heads instead of 3.
   Fix: run repro.py with smha_concat=True is set in config; change the default
     in SCITEConfig from smha_concat=False to smha_concat=True.

2. TRAINING EPOCHS — FIXED  [contributed ~3 pp; default now 200 with patience=30 early stopping]
   Paper trains up to 200 epochs and selects the best checkpoint by CV F1.
   Early stopping triggered at epoch ~58 on average (patience=30 on tri_f1).

3. FLAIR MODEL SUBSTITUTION — RESOLVED
   The original news-forward / news-backward models load correctly once flair.device
   is set to a proper torch.device instance (bug a below). The -fast substitution
   was unnecessary and has been reverted. The model now uses news-forward + news-backward
   as the paper specifies (each 2048-D stacked = 4096-D total input to BiLSTM).

4. TEST SET SIZE MISMATCH
   Paper: 804 sentences.  causalatee thagen/SCITE: 786 sentences.
   The 18 missing sentences may affect precision/recall comparability slightly.

5. EVALUATION PROTOCOL DIFFERENCE
   Paper: 5 independent runs × 200 epochs each, pick best checkpoint per run.
   Ours:  10-fold CV on train, 1 run, 5 epochs. Different in both run count and
   epoch budget, making direct number comparison imprecise.

6. FRAMEWORK
   Paper uses Keras 2.2.4.  Ours uses PyTorch.  Architecturally equivalent but
   weight initialisation defaults and gradient computation may differ slightly.

==============================================================
BUGS FIXED IN ORIGINAL CODE (not paper deviations)
==============================================================

a) flair.device set to torch.device CLASS not an instance
   → TypeError: device() argument 'type' must be str, not UntypedStorage
   Fix: flair.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

b) FlairEmbeddings("news-forward") appeared incompatible with PyTorch ≥ 2.0
   but was actually caused entirely by bug (a): once flair.device is a proper
   device instance the original models load and run correctly.
   Originally "fixed" by switching to -fast models; now reverted to news-forward /
   news-backward as the paper specifies.

c) config.flair_embedding_dim override is still kept in the code as a safeguard
   so that the actual stacked embedding length is always used regardless of which
   Flair model variant is loaded.

d) save_safetensors not disabled — Transformers rejects shared Flair weights
   Fix: save_safetensors=False in TrainingArguments

==============================================================
TODO / NEXT STEPS FOR FAITHFUL REPRODUCTION
==============================================================

1. ✓ smha_concat=True is now the default in SCITEConfig
2. ✓ --epochs default is now 200 with early stopping (patience=30)
3. ✓ LR scheduler (ReduceLROnPlateau) now wired via ReduceLROnPlateauCallback
4. Run 5 independent seeds, report mean ± std (currently 1 run)
5. Investigate the 18-sentence test set discrepancy vs paper's 804-sentence test
6. Investigate PyTorch-compatible loading of news-forward-0.4.1 to replace -fast models
