Skip to content

Fix GAMENet's Dynamic Memory - #1191

Open
lehendo wants to merge 1 commit into
sunlabuiuc:masterfrom
lehendo:gamenet-fix
Open

Fix GAMENet's Dynamic Memory#1191
lehendo wants to merge 1 commit into
sunlabuiuc:masterfrom
lehendo:gamenet-fix

Conversation

@lehendo

@lehendo lehendo commented Aug 19, 2026

Copy link
Copy Markdown
Collaborator

GAMENet.forward() hardcoded prev_drugs = torch.zeros(...) instead of the patient's actual drug history

require drugs_hist in the dataset schema, precompute a remap from drugs_hist's own input vocabulary to the drugs label_vocab used by ehr_adj/ddi_adj/the Memory Bank, and build the real multi-hot prev_drugs tensor from that remapped history

GAMENet.forward() hardcoded prev_drugs = torch.zeros(...) instead of the
patient's actual drug history, making the graph-augmented Dynamic Memory
mechanism the model is named for permanently inert. Per the paper (Shang
et al., "GAMENet: Graph Augmented MEmory Networks for Recommending
Medication Combination," AAAI 2019, arXiv:1809.01852), the Dynamic
Memory's values (Eq. 6) should be [c_m^1; ...; c_m^{t-1}] -- each
previous visit's actual administered drugs, retrieved via temporal
attention (Eq. 7). That attention/retrieval math was already correct;
only the memory's content was wrong.

pyhealth.tasks.drug_recommendation already produces exactly the needed
field (drugs_hist: nested per-visit drug history, current/target visit
zeroed out) -- GAMENet just never consumed it (an unused batch_to_multihot
import was a leftover sign of the abandoned wiring).

Fix: require drugs_hist in the dataset schema (a silent zeros-fallback
would just reintroduce the same silently-wrong-results bug in a new
form), precompute a remap from drugs_hist's own input vocabulary to the
drugs label_vocab used by ehr_adj/ddi_adj/the Memory Bank (they are
tokenized independently), and build the real multi-hot prev_drugs tensor
from that remapped history.

Verified end-to-end on real hardware: unit tests (including two new
regression tests), a 30-patient stress test with variable visit counts,
and the actual documented example script trained against real synthetic
MIMIC-III data.

@fbonc fbonc left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This problem predates your changes in this PR, but i think this implementation is incorrect for mixed-length batches now that dynamic memory is actually considered.

e.g., if patient A has 4 visits and patient B has 2, padding produces:
A: A1 A2 A3 A4
B: B1 B2 PAD PAD

queries[:, :-1] removes only the final column, producing:

A: A1 A2 A3
B: B1 B2 PAD

So this (correctly) excludes A’s current visit (A4), but it leaves B’s current visit (B2) AND padding in dynamic memory. This means softmax can assign attention to these invalid positions, taking weight away from B’s actual history (B1 here), which is bad.

I think to fix use the visit mask to exclude padding and each patient’s last valid visit before softmax. Also, patients without any prior visits should receive a zero dynamic-memory result.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants