DennisHuang648 commited on
Commit
80ed7c9
·
verified ·
1 Parent(s): 525dd7e

Upload gene_qformer_module.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. gene_qformer_module.py +191 -0
gene_qformer_module.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /data2/xiaoxinyu/project/model_merged_v75/gene_qformer_module.py
2
+ # -*- coding: utf-8 -*-
3
+
4
+ from __future__ import annotations
5
+ from dataclasses import dataclass
6
+ from typing import Optional
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+
11
+ try:
12
+ # transformers is optional for "read config only" behavior
13
+ from transformers import BertConfig, BertModel
14
+ except Exception:
15
+ BertConfig = None
16
+ BertModel = None
17
+
18
+
19
+ class _QFormerBlock(nn.Module):
20
+ """
21
+ A lightweight Q-Former block:
22
+ - self-attention on queries
23
+ - cross-attention: queries attend to gene tokens (kv)
24
+ - FFN
25
+ """
26
+
27
+ def __init__(self, dim: int, num_heads: int, dropout: float = 0.1):
28
+ super().__init__()
29
+ self.self_attn = nn.MultiheadAttention(
30
+ embed_dim=dim, num_heads=num_heads, dropout=dropout, batch_first=True
31
+ )
32
+ self.cross_attn = nn.MultiheadAttention(
33
+ embed_dim=dim, num_heads=num_heads, dropout=dropout, batch_first=True
34
+ )
35
+
36
+ self.norm_q1 = nn.LayerNorm(dim)
37
+ self.norm_q2 = nn.LayerNorm(dim)
38
+ self.norm_q3 = nn.LayerNorm(dim)
39
+
40
+ self.ffn = nn.Sequential(
41
+ nn.Linear(dim, dim * 4),
42
+ nn.GELU(),
43
+ nn.Dropout(dropout),
44
+ nn.Linear(dim * 4, dim),
45
+ nn.Dropout(dropout),
46
+ )
47
+
48
+ def forward(
49
+ self,
50
+ queries: torch.Tensor, # [B, Nq, D]
51
+ kv: torch.Tensor, # [B, L, D]
52
+ kv_key_padding_mask: Optional[torch.Tensor] = None, # [B, L], True for PAD
53
+ ) -> torch.Tensor:
54
+ # ---- Query self-attn ----
55
+ q = self.norm_q1(queries)
56
+ q2, _ = self.self_attn(q, q, q, need_weights=False)
57
+ queries = queries + q2
58
+
59
+ # ---- Cross-attn: queries attend to kv ----
60
+ q = self.norm_q2(queries)
61
+ q2, _ = self.cross_attn(
62
+ q, kv, kv,
63
+ key_padding_mask=kv_key_padding_mask, # True for PAD
64
+ need_weights=False,
65
+ )
66
+ queries = queries + q2
67
+
68
+ # ---- FFN ----
69
+ q = self.norm_q3(queries)
70
+ queries = queries + self.ffn(q)
71
+ return queries
72
+
73
+
74
+ class GeneQFormerBiomedBERT(nn.Module):
75
+ """
76
+ Gene Q-Former bridge module.
77
+
78
+ Why the name includes "BiomedBERT":
79
+ - Some papers initialize Q-Former from (BioMed)BERT.
80
+ - In your setup, you said you've merged BERT weights already, so
81
+ you can set load_pretrained_bert=False to avoid remote loading.
82
+ - We still optionally read the BERT config (hidden size, num layers, num heads)
83
+ to keep hyperparams consistent.
84
+
85
+ Inputs:
86
+ gene_tokens: [B, L, gene_in_dim] (e.g., 512)
87
+ gene_pad_mask: [B, L] bool, True indicates PAD positions (optional)
88
+
89
+ Output:
90
+ q_tokens: [B, num_queries, hidden] (e.g., [B,32,768])
91
+ """
92
+
93
+ def __init__(
94
+ self,
95
+ biomedbert_name: str = "",
96
+ gene_in_dim: int = 512,
97
+ hidden: int = 768,
98
+ num_queries: int = 32,
99
+ num_layers: int = 4,
100
+ num_heads: int = 12,
101
+ dropout: float = 0.1,
102
+ load_pretrained_bert: bool = False,
103
+ ):
104
+ super().__init__()
105
+
106
+ # Optionally read BERT config to align hyperparams
107
+ if biomedbert_name and BertConfig is not None:
108
+ try:
109
+ cfg = BertConfig.from_pretrained(biomedbert_name)
110
+ # Only override if caller didn't explicitly set hidden/layers/heads
111
+ # (We treat passed args as authoritative; config is a fallback.)
112
+ # Still, it's useful to sanity-check.
113
+ if hidden != cfg.hidden_size:
114
+ # keep user's hidden, but this can warn in logs if you want
115
+ pass
116
+ if num_heads != cfg.num_attention_heads:
117
+ pass
118
+ # if num_layers passed as default 4, but config has 12, you may want 12:
119
+ # We won't override automatically to avoid surprising behavior.
120
+ except Exception:
121
+ cfg = None
122
+ else:
123
+ cfg = None
124
+
125
+ self.gene_in_dim = int(gene_in_dim)
126
+ self.hidden = int(hidden)
127
+ self.num_queries = int(num_queries)
128
+ self.num_layers = int(num_layers)
129
+ self.num_heads = int(num_heads)
130
+
131
+ # Project gene token dim -> qformer hidden dim
132
+ self.gene_kv_proj = nn.Sequential(
133
+ nn.LayerNorm(self.gene_in_dim),
134
+ nn.Linear(self.gene_in_dim, self.hidden),
135
+ )
136
+
137
+ # Learnable query tokens
138
+ self.query_tokens = nn.Parameter(
139
+ torch.randn(1, self.num_queries, self.hidden) * 0.02
140
+ )
141
+
142
+ # Q-Former blocks
143
+ self.blocks = nn.ModuleList(
144
+ [_QFormerBlock(dim=self.hidden, num_heads=self.num_heads, dropout=dropout)
145
+ for _ in range(self.num_layers)]
146
+ )
147
+ self.out_norm = nn.LayerNorm(self.hidden)
148
+
149
+ # Optional: keep a BERTModel around (NOT used by default)
150
+ # If you later want to initialize weights from BERT, you can implement it here.
151
+ self._bert = None
152
+ if load_pretrained_bert:
153
+ if BertModel is None:
154
+ raise RuntimeError("transformers is not available, cannot load pretrained BERT.")
155
+ self._bert = BertModel.from_pretrained(biomedbert_name)
156
+
157
+ # NOTE: We do NOT directly plug BERT forward in this module,
158
+ # because BERT doesn't have cross-attn blocks by default.
159
+ # If you want to copy weights, implement a mapping routine
160
+ # (self-attn weights can be copied block-wise).
161
+ # For now, we just keep it loaded so you can manually inspect/copy.
162
+
163
+ def forward(
164
+ self,
165
+ gene_tokens: torch.Tensor, # [B, L, gene_in_dim]
166
+ gene_pad_mask: Optional[torch.Tensor] = None, # [B, L] bool, True for PAD
167
+ ) -> torch.Tensor:
168
+ if gene_tokens.dim() != 3:
169
+ raise ValueError(f"gene_tokens must be 3D [B,L,C], got {tuple(gene_tokens.shape)}")
170
+ B, L, C = gene_tokens.shape
171
+ if C != self.gene_in_dim:
172
+ raise ValueError(
173
+ f"gene_tokens last dim={C} != gene_in_dim={self.gene_in_dim}. "
174
+ f"Check Nicheformer output dim."
175
+ )
176
+
177
+ if gene_pad_mask is not None:
178
+ if gene_pad_mask.shape != (B, L):
179
+ raise ValueError(
180
+ f"gene_pad_mask shape {tuple(gene_pad_mask.shape)} != (B,L)=({B},{L})"
181
+ )
182
+ gene_pad_mask = gene_pad_mask.to(dtype=torch.bool, device=gene_tokens.device)
183
+
184
+ kv = self.gene_kv_proj(gene_tokens) # [B, L, hidden]
185
+
186
+ queries = self.query_tokens.expand(B, -1, -1).contiguous() # [B, Nq, hidden]
187
+
188
+ for blk in self.blocks:
189
+ queries = blk(queries, kv, kv_key_padding_mask=gene_pad_mask)
190
+
191
+ return self.out_norm(queries) # [B, Nq, hidden]