FlowKal commited on
Commit
4a7adf6
·
verified ·
1 Parent(s): e86d71f

Create model.py

Browse files
Files changed (1) hide show
  1. model.py +815 -0
model.py ADDED
@@ -0,0 +1,815 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import re
3
+ import math
4
+ import pickle
5
+ import random
6
+ from collections import defaultdict, namedtuple
7
+
8
+ import numpy as np
9
+ from PIL import Image
10
+ import torch
11
+ import torch.nn as nn
12
+ import torch.nn.functional as F
13
+ from torch.utils.data import Dataset, DataLoader
14
+ from torchvision import transforms, models
15
+ from tqdm import tqdm
16
+
17
+ # ================ HYPERPARAMETERS ================
18
+ IMG_HEIGHT = 256
19
+ IMG_WIDTH = 256
20
+ MAX_SEQ_LEN = 512
21
+ NUM_NESTED_LEVELS = 10
22
+ NUM_REL_POS = 3
23
+ PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN = '<PAD>', '<SOS>', '<EOS>', '<UNK>'
24
+ BATCH_SIZE = 8
25
+ NUM_EPOCHS = 300
26
+ MAX_LEARNING_RATE = 3e-4
27
+ WARMUP_RATIO = 0.1
28
+ WEIGHT_DECAY = 0.01
29
+ LAMBDA_POS = 0.2
30
+ LABEL_SMOOTHING = 0.0
31
+ DROPOUT_RATE = 0.3
32
+ NUM_LAYERS = 3
33
+ TEMPERATURE_INIT = 1.0
34
+ DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
35
+ SEED = 42
36
+ BEAM_SIZE = 10
37
+ MAX_GEN_LEN = MAX_SEQ_LEN
38
+
39
+ # ===== Fix seeds for reproducibility =====
40
+ random.seed(SEED)
41
+ np.random.seed(SEED)
42
+ torch.manual_seed(SEED)
43
+ if torch.cuda.is_available():
44
+ torch.cuda.manual_seed_all(SEED)
45
+
46
+ import re
47
+ from typing import List, Tuple, Union, Optional, Dict
48
+ import numpy as np
49
+
50
+
51
+ # Enhanced LaTeX tokenizer with better regex patterns
52
+ TOKEN_REGEX = re.compile(
53
+ r"(\\[a-zA-Z]+(?:\*)?)" # multi-letter commands (e.g. \frac, \sqrt, \begin*)
54
+ r"|(\\\{|\\\}|\\.)" # escaped chars and single-char commands
55
+ r"|([0-9]+(?:\.[0-9]+)?(?:[eE][+-]?[0-9]+)?)" # numbers (int, float, scientific)
56
+ r"|([A-Za-z]+)" # identifiers/variables
57
+ r"|([+\-*/=<>!~^_&|%])" # operators and structure markers
58
+ r"|([{}()\[\],.;:?'])" # delimiters and punctuation
59
+ r"|(\s+)" # whitespace (to handle properly)
60
+ r"|(\S)" # any other non-space char
61
+ )
62
+
63
+ def tokenize_latex(s: str) -> List[str]:
64
+ """
65
+ Split a LaTeX string into atomic tokens.
66
+ Filters out whitespace tokens.
67
+ """
68
+ # Исправление обработки специальных символов
69
+ s = re.sub(r'\\ ', ' ', s) # Обработка пробелов
70
+ s = re.sub(r'\\\n', '', s) # Удаление переносов
71
+
72
+ tokens = []
73
+ for match in TOKEN_REGEX.finditer(s):
74
+ token = match.group(0)
75
+ if token.isspace():
76
+ continue
77
+ # Это необязательная, но потенциально полезная эвристика для \left{ и \right}
78
+ # Если она вызывает проблемы, можно убрать.
79
+ if token in {'{', '}'} and tokens and tokens[-1] in {'\\left', '\\right'}:
80
+ tokens[-1] += token
81
+ else:
82
+ tokens.append(token)
83
+
84
+ return tokens
85
+
86
+ # Position Forest Implementation
87
+ STRUCTURE_CMDS = {'^', '_', '\\sqrt', '\\frac', '\\sum', '\\int', '\\lim'}
88
+
89
+ class PositionForestEncoder:
90
+ def position_forest_ids(self, tokens):
91
+ cache_key = tuple(tokens)
92
+ if hasattr(self, '_cache'):
93
+ if cache_key in self._cache:
94
+ return self._cache[cache_key]
95
+ else:
96
+ self._cache = {}
97
+ T = len(tokens)
98
+ ids = ['M']*T
99
+ def find_close(i):
100
+ depth = 0
101
+ for j in range(i, T):
102
+ if tokens[j]=='{': depth+=1
103
+ elif tokens[j]=='}':
104
+ depth-=1
105
+ if depth==0: return j
106
+ return T-1
107
+ def mark(l,r,label):
108
+ for k in range(max(0,l), min(T,r+1)):
109
+ ids[k]+=label
110
+ def rec(l,r):
111
+ i = l
112
+ while i<=r:
113
+ tok = tokens[i]
114
+ # handle ^, _, ...
115
+ if tok in ('^','_') and i+1<=r:
116
+ if tokens[i+1]=='{':
117
+ e=find_close(i+1)
118
+ lbl='L' if tok=='^' else 'R'
119
+ mark(i+2,e-1,lbl)
120
+ rec(i+2,e-1)
121
+ i=e+1
122
+ else:
123
+ lbl='L' if tok=='^' else 'R'
124
+ mark(i+1,i+1,lbl)
125
+ i+=2
126
+ elif tok=='\\sqrt' and i+1<=r and tokens[i+1]=='{':
127
+ e=find_close(i+1)
128
+ mark(i+2,e-1,'L'); rec(i+2,e-1)
129
+ i=e+1
130
+ elif tok=='\\frac' and i+1<=r and tokens[i+1]=='{':
131
+ n_end=find_close(i+1)
132
+ if n_end+1<=r and tokens[n_end+1]=='{':
133
+ d_end=find_close(n_end+1)
134
+ mark(i+2,n_end-1,'L'); rec(i+2,n_end-1)
135
+ mark(n_end+2,d_end-1,'R'); rec(n_end+2,d_end-1)
136
+ i=d_end+1
137
+ else: i+=1
138
+ else: i+=1
139
+ rec(0,T-1)
140
+ self._cache[cache_key] = ids
141
+ return ids
142
+
143
+ def forest_MLR(self, latex: str):
144
+ tokens = tokenize_latex(latex)
145
+ if not tokens: tokens=['']
146
+ pos_ids = self.position_forest_ids(tokens)
147
+ nested_depth = [len(pid)-1 for pid in pos_ids]
148
+ rel_pos = [1 if pid.endswith('L') else 2 if pid.endswith('R') else 0 for pid in pos_ids]
149
+ # add special tokens
150
+ tokens = [SOS_TOKEN]+tokens+[EOS_TOKEN]
151
+ nested_depth = [0]+nested_depth+[0]
152
+ rel_pos = [0]+rel_pos+[0]
153
+ # pad/truncate
154
+ curr=len(tokens)
155
+ if curr>MAX_SEQ_LEN:
156
+ tokens=tokens[:MAX_SEQ_LEN]; nested_depth=nested_depth[:MAX_SEQ_LEN]; rel_pos=rel_pos[:MAX_SEQ_LEN]; curr=MAX_SEQ_LEN
157
+ pad_len=MAX_SEQ_LEN-curr
158
+ tokens+= [PAD_TOKEN]*pad_len
159
+ nested_depth+= [0]*pad_len; rel_pos+=[0]*pad_len
160
+ attention_mask = [1]*curr + [0]*pad_len
161
+ return tokens, np.array(nested_depth), np.array(rel_pos), np.array(attention_mask)
162
+
163
+ # ===== Vocabulary =====
164
+ class Vocab:
165
+ def __init__(self, expressions):
166
+ freq=defaultdict(int)
167
+ for expr in expressions:
168
+ for t in tokenize_latex(expr): freq[t]+=1
169
+ self.itos=[PAD_TOKEN,SOS_TOKEN,EOS_TOKEN,UNK_TOKEN]+sorted(freq.keys(), key=lambda x:-freq[x])
170
+ self.stoi={tok:i for i,tok in enumerate(self.itos)}
171
+ def encode(self, tokens):
172
+ return [self.stoi.get(t,self.stoi[UNK_TOKEN]) for t in tokens]
173
+ def decode(self, ids):
174
+ return ''.join(self.itos[i] for i in ids if self.itos[i] not in (PAD_TOKEN,SOS_TOKEN,EOS_TOKEN))
175
+
176
+ # ===== Dataset =====
177
+ class CROHMEDataset(Dataset):
178
+ def __init__(self, base_dir, caption_file, vocab, pfe):
179
+ self.vocab=vocab; self.pfe=pfe; self.items=[]
180
+ img_dir=os.path.join(base_dir,'img')
181
+ with open(os.path.join(base_dir,caption_file),'r',encoding='utf-8') as f:
182
+ for line in f:
183
+ fid,lt=line.strip().split('\t',1)
184
+ self.items.append((os.path.join(img_dir,fid+'.bmp'),lt))
185
+ self.transform = transforms.Compose([
186
+ transforms.Resize((IMG_HEIGHT, IMG_WIDTH)),
187
+ transforms.RandomAffine(degrees=5, translate=(0.05, 0.05), scale=(0.95, 1.05), shear=5),
188
+ transforms.ColorJitter(brightness=0.3, contrast=0.3),
189
+ transforms.ToTensor(),
190
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
191
+ ])
192
+ def __len__(self): return len(self.items)
193
+ def __getitem__(self,idx):
194
+ path,latex=self.items[idx]
195
+ img=Image.open(path).convert('RGB'); img=self.transform(img)
196
+ tokens,nested,rel,mask=self.pfe.forest_MLR(latex)
197
+ token_ids=self.vocab.encode(tokens)
198
+ return img, {
199
+ 'token_ids': torch.tensor(token_ids,dtype=torch.long),
200
+ 'nested': torch.tensor(nested,dtype=torch.long),
201
+ 'relpos': torch.tensor(rel,dtype=torch.long),
202
+ 'attention_mask': torch.tensor(mask,dtype=torch.long)
203
+ }, latex
204
+
205
+
206
+ # ===== Improved Implicit Attention Correction =====
207
+ class IAC(nn.Module):
208
+ """
209
+ Implicit Attention Correction модуль согласно статье
210
+ Исправленная версия с правильной логикой накопления внимания
211
+ """
212
+ def __init__(self, d_model=256, num_heads=8, dropout=DROPOUT_RATE, structure_symbols_set=None,vocab=None):
213
+ super().__init__()
214
+ assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
215
+ self.d_model = d_model
216
+ self.num_heads = num_heads
217
+ self.head_dim = d_model // num_heads
218
+
219
+ # Структурные символы - те, которые НЕ имеют визуального представления
220
+ self.structure_symbols = structure_symbols_set or {
221
+ '^', '{', '}', '_', EOS_TOKEN,UNK_TOKEN,SOS_TOKEN,PAD_TOKEN
222
+ }
223
+
224
+ # Linear projections для Q, K, V
225
+ self.wq = nn.Linear(d_model, d_model, bias=False)
226
+ self.wk = nn.Linear(d_model, d_model, bias=False)
227
+ self.wv = nn.Linear(d_model, d_model, bias=False)
228
+ self.wo = nn.Linear(d_model, d_model)
229
+
230
+ # φ функция для обработки accumulated attention
231
+ # Используем простую архитектуру: conv + linear
232
+ self.phi_conv = nn.Conv2d(num_heads, num_heads, kernel_size=3, padding=1, groups=num_heads)
233
+ self.phi_linear = nn.Linear(num_heads, d_model)
234
+ self.phi_norm = nn.LayerNorm(d_model)
235
+
236
+ self.attn_dropout = nn.Dropout(dropout)
237
+ self.proj_dropout = nn.Dropout(dropout)
238
+
239
+ # Накопленное внимание для коррекции
240
+ self.register_buffer('accumulated_attention', None, persistent=False)
241
+ structure_ids_list = [vocab.stoi.get(s, -1) for s in structure_symbols_set]
242
+ self.register_buffer('structure_ids',
243
+ torch.tensor([sid for sid in structure_ids_list if sid != -1], dtype=torch.long))
244
+ self._init_weights()
245
+
246
+
247
+ def _init_weights(self):
248
+ for m in [self.wq, self.wk, self.wv, self.wo, self.phi_linear]:
249
+ nn.init.xavier_uniform_(m.weight)
250
+ if m.bias is not None:
251
+ nn.init.constant_(m.bias, 0)
252
+
253
+ nn.init.xavier_uniform_(self.phi_conv.weight)
254
+ nn.init.constant_(self.phi_conv.bias, 0)
255
+
256
+ def reset_state(self):
257
+ """Сброс состояния IAC - важно для начала новой последовательности"""
258
+ self.accumulated_attention = None
259
+
260
+ def _is_structure_symbol(self, symbol):
261
+ """
262
+ Проверяет, является ли символ структурным
263
+ Структурные символы не имеют визуального представления в изображении
264
+ """
265
+ if isinstance(symbol, (list, tuple)):
266
+ return [sym in self.structure_symbols for sym in symbol]
267
+ return symbol in self.structure_symbols
268
+
269
+ def _compute_phi(self, accumulated_attention):
270
+ """
271
+ Вычисляет φ(A^k) - функцию коррекции на основе накопленного внимания.
272
+ ИСПРАВЛЕННАЯ ВЕРСИЯ для обработки 5D тензора.
273
+
274
+ Args:
275
+ accumulated_attention: [B, H, Lq, Hp, Wp] - история накопленного внимания
276
+ """
277
+ # 1. Получаем 5 измерений
278
+ B, H, Lq, Hp, Wp = accumulated_attention.shape
279
+
280
+ # 2. Объединяем B и Lq в одно измерение, чтобы подать в Conv2d.
281
+ # Conv2d ожидает на вход (N, C_in, H_in, W_in).
282
+ # Наша C_in - это H (количество голов).
283
+ # (B, H, Lq, Hp, Wp) -> (B, Lq, H, Hp, Wp) -> (B * Lq, H, Hp, Wp)
284
+ x = accumulated_attention.permute(0, 2, 1, 3, 4).contiguous()
285
+ x = x.view(B * Lq, H, Hp, Wp)
286
+
287
+ # 3. Применяем conv2d для извлечения локальных паттернов покрытия
288
+ phi_conv_out = self.phi_conv(x) # Shape: [B * Lq, H, Hp, Wp]
289
+
290
+ # 4. Преобразуем для подачи в Linear слой.
291
+ # (B*Lq, H, Hp, Wp) -> (B*Lq, Hp, Wp, H) -> (B*Lq, Hp*Wp, H)
292
+ phi_permuted = phi_conv_out.permute(0, 2, 3, 1)
293
+ phi_flat = phi_permuted.contiguous().view(B * Lq, Hp * Wp, H)
294
+
295
+ # 5. Проецируем в d_model пространство и нормализуем
296
+ phi_features = self.phi_linear(phi_flat) # Shape: [B * Lq, Hp*Wp, d_model]
297
+ phi_features = self.phi_norm(phi_features)
298
+
299
+ # 6. Возвращаем обратно в форму, совместимую с forward pass.
300
+ # Разделяем B и Lq обратно.
301
+ # (B*Lq, Hp*Wp, d_model) -> (B, Lq, Hp*Wp, d_model)
302
+ phi_out = phi_features.view(B, Lq, Hp * Wp, self.d_model)
303
+
304
+ return phi_out
305
+
306
+ def forward(self, q, k, v, current_symbols_ids=None, key_padding_mask=None):
307
+ """
308
+ Forward pass с применением IAC коррекции (батчево для всей последовательности)
309
+ Args:
310
+ q: [B, Lq, d_model]
311
+ k: [B, Lk, d_model]
312
+ v: [B, Lv, d_model]
313
+ current_symbols: [B, Lq] (для обучения) или List[str] (для инференса)
314
+ key_padding_mask: [B, Lk]
315
+ """
316
+ B, Lq, _ = q.size()
317
+ _, Lk, _ = k.size()
318
+ D = self.head_dim
319
+ spatial_size = int(math.sqrt(Lk))
320
+ assert spatial_size * spatial_size == Lk, f"Feature map должна быть квадратной, получено {Lk}"
321
+
322
+ Q = self.wq(q).view(B, Lq, self.num_heads, D).transpose(1, 2) # [B, H, Lq, D]
323
+ K = self.wk(k).view(B, Lk, self.num_heads, D).transpose(1, 2) # [B, H, Lk, D]
324
+ V = self.wv(v).view(B, Lk, self.num_heads, D).transpose(1, 2) # [B, H, Lk, D]
325
+
326
+ scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(D) # [B, H, Lq, Lk]
327
+ if key_padding_mask is not None:
328
+ mask = key_padding_mask.unsqueeze(1).unsqueeze(2)
329
+ scores = scores.masked_fill(mask, float('-inf'))
330
+ attn_weights = F.softmax(scores, dim=-1) # [B, H, Lq, Lk]
331
+
332
+ # Батчевая коррекция внимания для всего Lq (только при обучении)
333
+ if current_symbols_ids is not None and self.accumulated_attention is not None:
334
+ # Проверяем совместимость размеров
335
+ accum_B, accum_H, accum_Lq, accum_Hp, accum_Wp = self.accumulated_attention.shape
336
+
337
+ # Если размеры не совпадают, расширяем/обрезаем накопленное внимание
338
+ if accum_Lq < Lq:
339
+ # Дополняем нулями
340
+ padding = torch.zeros(
341
+ accum_B, accum_H, Lq - accum_Lq, accum_Hp, accum_Wp,
342
+ device=self.accumulated_attention.device
343
+ )
344
+ accumulated_attention = torch.cat([self.accumulated_attention, padding], dim=2)
345
+ elif accum_Lq > Lq:
346
+ # Берем только первые Lq элементов
347
+ accumulated_attention = self.accumulated_attention[:, :, :Lq, :, :]
348
+ else:
349
+ accumulated_attention = self.accumulated_attention
350
+
351
+ phi_correction = self._compute_phi(accumulated_attention)
352
+ # Преобразуем phi_correction к attention space
353
+ phi_attn = phi_correction.view(B, Lq, spatial_size * spatial_size, self.d_model)
354
+ phi_attn = phi_attn.view(B, Lq, spatial_size * spatial_size, self.num_heads, self.head_dim).permute(0, 3, 1, 4, 2) # [B, H, Lq, D, Lk]
355
+ # Q: [B, H, Lq, D], phi_attn: [B, H, Lq, D, Lk]
356
+ phi_scores = torch.matmul(Q.unsqueeze(-2), phi_attn).squeeze(-2) # [B, H, Lq, Lk]
357
+ corrected_scores = scores - phi_scores
358
+ attn_weights = F.softmax(corrected_scores, dim=-1)
359
+
360
+ # Обновляем накопленное внимание ПОСЛЕ коррекции
361
+ if current_symbols_ids is not None:
362
+ self._update_accumulated_attention(attn_weights, current_symbols_ids, spatial_size)
363
+
364
+ attn_weights = self.attn_dropout(attn_weights)
365
+ context = torch.matmul(attn_weights, V) # [B, H, Lq, D]
366
+ context = context.transpose(1, 2).reshape(B, Lq, self.d_model)
367
+ output = self.wo(context)
368
+ output = self.proj_dropout(output)
369
+ return output, attn_weights
370
+
371
+
372
+ def _update_accumulated_attention(self, attn_weights, current_symbols_ids, spatial_size):
373
+ """
374
+ Векторизованная версия без медленных циклов Python.
375
+ """
376
+ B, H, Lq, Lk = attn_weights.shape
377
+
378
+ # Создаем маску за одну быструю операцию на GPU
379
+ # Сравниваем каждый ID в батче со списком ID структурных символов
380
+ is_structure_mask = (current_symbols_ids.unsqueeze(-1) == self.structure_ids.view(1, 1, -1)).any(dim=-1)
381
+
382
+ # Инвертируем маску (1.0 для обычных символов, 0.0 для структурных)
383
+ # и приводим к нужному виду для умножения
384
+ indicator = (~is_structure_mask).float().view(B, 1, Lq, 1)
385
+
386
+ # Дальнейшая логика остается той же, но теперь она работает на GPU без тормозов
387
+ masked_attn = attn_weights * indicator
388
+
389
+ if Lq > 1:
390
+ masked_attn_shifted = torch.zeros_like(masked_attn)
391
+ masked_attn_shifted[:, :, 1:] = masked_attn[:, :, :-1]
392
+ accumulated = torch.cumsum(masked_attn_shifted, dim=2)
393
+ else:
394
+ accumulated = torch.zeros_like(masked_attn)
395
+
396
+ self.accumulated_attention = accumulated.view(B, H, Lq, spatial_size, spatial_size)
397
+
398
+
399
+ class EnhancedDecoderLayer(nn.Module):
400
+ """
401
+ Decoder layer с поддержкой IAC - исправленная версия
402
+ """
403
+ def __init__(self, d_model=256, num_heads=8, d_ff=1024, dropout=DROPOUT_RATE, structure_symbols_set=None,vocab=None):
404
+ super().__init__()
405
+ self.d_model = d_model
406
+
407
+ # Self-attention
408
+ self.norm1 = nn.LayerNorm(d_model)
409
+ self.self_attn = nn.MultiheadAttention(d_model, num_heads, dropout=dropout, batch_first=True)
410
+
411
+ # Cross-attention with IAC
412
+ self.norm2 = nn.LayerNorm(d_model)
413
+ self.iac = IAC(d_model, num_heads, dropout, structure_symbols_set,vocab)
414
+
415
+ # Feed-forward
416
+ self.norm3 = nn.LayerNorm(d_model)
417
+ self.ffn = nn.Sequential(
418
+ nn.Linear(d_model, d_ff),
419
+ nn.GELU(),
420
+ nn.Dropout(dropout),
421
+ nn.Linear(d_ff, d_model),
422
+ nn.Dropout(dropout)
423
+ )
424
+
425
+ self.dropout = nn.Dropout(dropout)
426
+
427
+ def reset_iac_state(self):
428
+ """Сброс состояния IAC"""
429
+ self.iac.reset_state()
430
+
431
+ def forward(self, x, encoder_output, current_symbols_ids=None,
432
+ tgt_mask=None, tgt_key_padding_mask=None, memory_key_padding_mask=None):
433
+ """
434
+ Forward pass decoder layer
435
+
436
+ Args:
437
+ x: [B, Lq, d_model] - target embeddings
438
+ encoder_output: [B, Lk, d_model] - encoder output (visual features)
439
+ current_symbols: str или List[str] - текущие символы для IAC
440
+ tgt_mask: causal mask для self-attention
441
+ tgt_key_padding_mask: padding mask для target
442
+ memory_key_padding_mask: padding mask для encoder output
443
+ """
444
+ batch_size = x.size(0)
445
+
446
+ # Self-attention
447
+ residual = x
448
+ x = self.norm1(x)
449
+ self_attn_output, self_attn_weights = self.self_attn(
450
+ x, x, x,
451
+ attn_mask=tgt_mask,
452
+ key_padding_mask=tgt_key_padding_mask
453
+ )
454
+ x = residual + self.dropout(self_attn_output)
455
+
456
+ # Cross-attention с IAC
457
+ residual = x
458
+ x = self.norm2(x)
459
+
460
+ # IAC применяется только к cross-attention
461
+ cross_attn_output, cross_attn_weights = self.iac(
462
+ x, encoder_output, encoder_output,
463
+ current_symbols_ids=current_symbols_ids,
464
+ key_padding_mask=memory_key_padding_mask
465
+ )
466
+ x = residual + self.dropout(cross_attn_output)
467
+
468
+ # Feed-forward
469
+ residual = x
470
+ x = self.norm3(x)
471
+ ffn_output = self.ffn(x)
472
+ x = residual + self.dropout(ffn_output)
473
+
474
+ return x
475
+
476
+ class PosFormerImprovedWithIAC(nn.Module):
477
+ """
478
+ PosFormer с исправленной поддержкой IAC и правильной реализацией согласно статье
479
+ """
480
+ def __init__(
481
+ self,
482
+ vocab_size,
483
+ vocab,
484
+ pad_token_id,
485
+ structure_symbols_set,
486
+ id_to_token_map,
487
+ d_model=256,
488
+ num_heads=8,
489
+ num_layers=NUM_LAYERS,
490
+ d_ff=1024,
491
+ dropout=0.1,
492
+ max_seq_len=MAX_SEQ_LEN,
493
+ max_nested_levels=NUM_NESTED_LEVELS,
494
+ identifier_vocab_size=None,
495
+ identifier_max_len=10,
496
+ ):
497
+ super().__init__()
498
+ self.vocab_size = vocab_size
499
+ self.vocab = vocab
500
+ self.pad_id = pad_token_id
501
+ self.structure_symbols = structure_symbols_set
502
+ self.id_to_token = id_to_token_map
503
+ self.max_seq_len = max_seq_len
504
+ self.max_nested_levels = max_nested_levels
505
+ self.identifier_max_len = identifier_max_len
506
+ self.d_model = d_model
507
+
508
+ # Encoder
509
+ self.enc = ImprovedImageEncoder(d_model)
510
+
511
+ # Token embeddings + pos
512
+ self.emb_tok = nn.Embedding(vocab_size, d_model)
513
+ self.pos_enc = nn.Parameter(torch.randn(1, max_seq_len, d_model) * 0.02)
514
+ self.norm_in = nn.LayerNorm(d_model)
515
+
516
+ # Identifier embeddings (ξ function)
517
+ if identifier_vocab_size is not None:
518
+ self.identifier_emb = nn.Embedding(identifier_vocab_size, d_model)
519
+ self.xi_function = nn.Sequential(
520
+ nn.Linear(d_model, d_model),
521
+ nn.GELU(),
522
+ nn.LayerNorm(d_model)
523
+ )
524
+ else:
525
+ self.identifier_emb = None
526
+ self.xi_function = None
527
+ self.identifier_pos_enc = nn.Parameter(torch.randn(1, max_seq_len, d_model) * 0.02)
528
+
529
+ # Decoder layers с IAC
530
+ self.decoders = nn.ModuleList([
531
+ EnhancedDecoderLayer(d_model, num_heads, d_ff, dropout, structure_symbols_set,vocab=vocab)
532
+ for _ in range(num_layers)
533
+ ])
534
+ self.norm_out = nn.LayerNorm(d_model)
535
+
536
+ # Prediction heads
537
+ self.W_n = nn.Linear(d_model, max_nested_levels + 1)
538
+ self.W_r = nn.Linear(d_model, 3) # M, L, R
539
+ self.head_tok = nn.Linear(d_model, vocab_size)
540
+
541
+ self._init_weights()
542
+
543
+ def _init_weights(self):
544
+ """Правильная инициализация весов"""
545
+ # Embedding layers
546
+ nn.init.normal_(self.emb_tok.weight, std=0.02)
547
+ if self.identifier_emb is not None:
548
+ nn.init.normal_(self.identifier_emb.weight, std=0.02)
549
+
550
+ # Position encodings
551
+ nn.init.normal_(self.pos_enc, std=0.02)
552
+ nn.init.normal_(self.identifier_pos_enc, std=0.02)
553
+
554
+ # Prediction heads
555
+ nn.init.xavier_uniform_(self.W_n.weight)
556
+ nn.init.constant_(self.W_n.bias, 0.0)
557
+ nn.init.xavier_uniform_(self.W_r.weight)
558
+ nn.init.constant_(self.W_r.bias, 0.0)
559
+ nn.init.xavier_uniform_(self.head_tok.weight)
560
+ nn.init.constant_(self.head_tok.bias, 0.0)
561
+
562
+ def reset_iac_state(self):
563
+ """Сброс состояния IAC для всех decoder layers"""
564
+ for decoder in self.decoders:
565
+ decoder.reset_iac_state()
566
+
567
+ def process_identifiers(self, identifiers):
568
+ """
569
+ Обработка identifier embeddings согласно формуле (1)
570
+ Q_emb = [ξ(Q_1); ξ(Q_2); ...; ξ(Q_L)] + Q_pos
571
+ """
572
+ if self.identifier_emb is None:
573
+ return None
574
+
575
+ B, T, U = identifiers.size()
576
+
577
+ # Embedding lookup
578
+ emb = self.identifier_emb(identifiers) # [B, T, U, d_model]
579
+
580
+ # Маска для padding
581
+ pad_mask = identifiers == 0
582
+ mask = ~pad_mask.unsqueeze(-1) # [B, T, U, 1]
583
+
584
+ # Применяем маску и усредняем
585
+ emb = emb * mask.float()
586
+ lengths = mask.sum(dim=2).float().clamp(min=1) # [B, T, 1]
587
+ avg_emb = emb.sum(dim=2) / lengths # [B, T, d_model]
588
+
589
+ # Применяем ξ function
590
+ processed = self.xi_function(avg_emb)
591
+
592
+ # Добавляем позиционное кодирование
593
+ return processed + self.identifier_pos_enc[:, :T, :]
594
+
595
+ def get_symbol_from_id(self, tid):
596
+ """Получение символа по ID"""
597
+ if isinstance(tid, torch.Tensor):
598
+ tid = tid.item()
599
+ return self.id_to_token.get(tid, f"<unk_{tid}>")
600
+
601
+ def forward(self, imgs, feat=None, mode='train'):
602
+ """Основная функция forward"""
603
+ if mode == 'train':
604
+ return self._forward_train(imgs, feat)
605
+ else:
606
+ return self._forward_inference(imgs)
607
+
608
+ def _forward_train(self, imgs, decoder_input):
609
+ """
610
+ Принимает сдвинутый decoder_input и возвращает логиты для всех голов.
611
+ """
612
+ B, T = decoder_input.size()
613
+
614
+ self.reset_iac_state()
615
+ enc = self.enc(imgs)
616
+
617
+ x = self.emb_tok(decoder_input)
618
+ x = x + self.pos_enc[:, :T, :]
619
+ x = self.norm_in(x)
620
+
621
+ tgt_mask = torch.triu(torch.ones(T, T, device=imgs.device), diagonal=1).bool()
622
+ pad_mask = decoder_input == self.pad_id
623
+
624
+
625
+ for decoder in self.decoders:
626
+ x = decoder(
627
+ x, enc,
628
+ current_symbols_ids=decoder_input, # Передаем ID напрямую
629
+ tgt_mask=tgt_mask,
630
+ tgt_key_padding_mask=pad_mask
631
+ )
632
+
633
+ x = self.norm_out(x)
634
+
635
+ # Получаем логиты от всех трех "голов"
636
+ nested_logits = self.W_n(x)
637
+ relpos_logits = self.W_r(x)
638
+ token_logits = self.head_tok(x)
639
+
640
+ return {
641
+ 'token_ids': token_logits,
642
+ 'nested': nested_logits,
643
+ 'relpos': relpos_logits,
644
+ 'features': x
645
+ }
646
+
647
+ def forward_inference(self, enc_output, tokens): # Убираем current_symbols
648
+ """
649
+ Инференс с заданными энкодером и токенами.
650
+ Теперь передает ID токенов напрямую в декодер.
651
+ """
652
+ B, T = tokens.size()
653
+
654
+ # Эмбеддинги и позиционное кодирование
655
+ x = self.emb_tok(tokens) + self.pos_enc[:, :T, :]
656
+ x = self.norm_in(x)
657
+
658
+ # Прямой проход через декодер
659
+ for decoder in self.decoders:
660
+ x = decoder( # Возвращаемые веса нам здесь не нужны
661
+ x,
662
+ enc_output,
663
+ current_symbols_ids=tokens, # Передаем тензор с ID
664
+ tgt_mask=None # При инференсе каузальная маска не нужна
665
+ )
666
+
667
+ x = self.norm_out(x)
668
+ token_logits = self.head_tok(x)
669
+
670
+ return {'logits': token_logits}
671
+
672
+ # ===== Image Encoder =====
673
+ class ImprovedImageEncoder(nn.Module):
674
+ def __init__(self, d_model=256):
675
+ super().__init__()
676
+ # Загрузка предобученной DenseNet-121
677
+ base = models.densenet121(pretrained=True)
678
+
679
+ # Используем все слои до последнего пулинга
680
+ self.features = base.features
681
+
682
+ # Проекция в d_model
683
+ self.spatial_proj = nn.Conv2d(1024, d_model, 1)
684
+
685
+ # Адаптивный пулинг для фиксированного размера
686
+ self.adaptive_pool = nn.AdaptiveAvgPool2d((7, 7)) # Фиксируем размер 7x7
687
+
688
+ # Позиционные эмбеддинги
689
+ self.pos_embed = nn.Parameter(torch.randn(1, 49, d_model) * 0.1)
690
+ self.norm = nn.LayerNorm(d_model)
691
+
692
+ # Dropout для регуляризации
693
+ self.dropout = nn.Dropout(0.1)
694
+
695
+ def forward(self, x):
696
+ # Прямой проход через DenseNet
697
+ f = self.features(x) # [B, 1024, H, W]
698
+
699
+ # Проекция в пространство d_model
700
+ s = self.spatial_proj(f) # [B, d_model, H, W]
701
+
702
+ # Адаптивный пулинг для фиксированного размера
703
+ s = self.adaptive_pool(s) # [B, d_model, 7, 7]
704
+
705
+ # Преобразование: [B, C, H, W] -> [B, H*W, C]
706
+ B, C, H, W = s.shape
707
+ s = s.flatten(2).transpose(1, 2) # [B, 49, d_model]
708
+
709
+ # Добавление позиционных эмбеддингов
710
+ s = s + self.pos_embed
711
+
712
+ # Нормализация и dropout
713
+ s = self.norm(s)
714
+ s = self.dropout(s)
715
+
716
+ return s
717
+ def generate_beam_search(model, image_tensor, beam_size, max_len, sos_id, eos_id, pad_id, device):
718
+ """
719
+ Генерирует последовательность токенов с использованием Beam Search.
720
+ """
721
+ model.eval()
722
+
723
+ model.reset_iac_state()
724
+
725
+ # 2. Получаем признаки из энкодера один раз
726
+ with torch.no_grad():
727
+ # enc_output всегда имеет batch_size=1
728
+ enc_output = model.enc(image_tensor.unsqueeze(0))
729
+
730
+ # 3. Инициализация лучей.
731
+ # beams - это список из кортежей (последовательность_тензор, score)
732
+ beams = [(torch.tensor([sos_id], dtype=torch.long, device=device), 0.0)]
733
+
734
+ # 4. Пошаговая генерация
735
+ for t in range(1, max_len):
736
+ all_candidates = []
737
+
738
+ for seq, score in beams:
739
+ if seq[-1] == eos_id:
740
+ all_candidates.append((seq, score))
741
+ continue
742
+
743
+ with torch.no_grad():
744
+ input_seq = seq.unsqueeze(0)
745
+
746
+
747
+ outputs = model.forward_inference(enc_output, input_seq)
748
+
749
+ logits = outputs['logits'][:, -1, :] # Логиты для последнего токена
750
+
751
+ log_probs = F.log_softmax(logits, dim=-1)
752
+ top_log_probs, top_ids = torch.topk(log_probs, beam_size, dim=-1)
753
+
754
+ # Создаем новых кандидатов
755
+ for i in range(beam_size):
756
+ next_id = top_ids[0, i]
757
+ log_prob = top_log_probs[0, i].item()
758
+
759
+ # Создаем новую последовательность, добавляя новый токен
760
+ new_seq = torch.cat([seq, next_id.view(1)])
761
+ new_score = score + log_prob
762
+ all_candidates.append((new_seq, new_score))
763
+
764
+ # 5. Сортируем всех кандидатов и выбираем `beam_size` лучших
765
+ ordered = sorted(all_candidates, key=lambda x: x[1], reverse=True)
766
+ beams = ordered[:beam_size]
767
+
768
+ # 6. Условие остановки: если все лучшие лучи закончились на EOS
769
+ if all(b[0][-1] == eos_id for b in beams):
770
+ break
771
+
772
+ # 7. Выбираем лучший луч, нормализуя на длину, чтобы не штрафовать длинные последовательности
773
+ best_beam = sorted(beams, key=lambda x: x[1] / len(x[0]), reverse=True)[0]
774
+ best_seq = best_beam[0]
775
+
776
+ return best_seq
777
+
778
+ # ======================================================================
779
+ # 2. Удобная обертка для предсказания
780
+ # ======================================================================
781
+
782
+ def predict_with_beam_search(model, image_path, vocab, beam_size, device):
783
+ """
784
+ Загружает изображение, запускает Beam Search и возвращает LaTeX строку.
785
+ """
786
+ # 1. Загрузка и трансформация изображения
787
+ transform = transforms.Compose([
788
+ transforms.Resize((IMG_HEIGHT, IMG_WIDTH)),
789
+ transforms.ToTensor(),
790
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
791
+ ])
792
+ image = Image.open(image_path).convert("RGB")
793
+ image = ImageOps.invert(image)
794
+ image_tensor = transform(image).to(device)
795
+
796
+ # 2. Получение ID спецтокенов
797
+ sos_id = vocab.stoi[SOS_TOKEN]
798
+ eos_id = vocab.stoi[EOS_TOKEN]
799
+ pad_id = vocab.stoi[PAD_TOKEN]
800
+
801
+ # 3. Генерация
802
+ predicted_ids = generate_beam_search(
803
+ model, image_tensor, beam_size, MAX_SEQ_LEN, sos_id, eos_id, pad_id, device
804
+ )
805
+
806
+ # 4. Декодирование
807
+ predicted_ids = predicted_ids.cpu().numpy()
808
+ try:
809
+ eos_index = list(predicted_ids).index(eos_id)
810
+ predicted_ids = predicted_ids[:eos_index]
811
+ except ValueError:
812
+ pass # EOS не был сгенерирован
813
+
814
+ latex_string = vocab.decode(predicted_ids)
815
+ return latex_string