FlowKal commited on
Commit
67f4bb3
·
verified ·
1 Parent(s): 5a758a6

Update model.py

Browse files
Files changed (1) hide show
  1. model.py +650 -740
model.py CHANGED
@@ -1,798 +1,708 @@
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
- """
722
- model.eval()
723
- model.reset_iac_state()
724
-
725
- # 2. Получаем признаки из энкодера один раз.
726
- # image_tensor уже должен иметь batch_size=1, поэтому убираем .unsqueeze(0)
727
- with torch.no_grad():
728
- enc_output = model.enc(image_tensor)
729
-
730
- # 3. Инициализация лучей.
731
- beams = [(torch.tensor([sos_id], dtype=torch.long, device=device), 0.0)]
732
-
733
- # 4. Пошаговая генерация
734
- for t in range(1, max_len):
735
- all_candidates = []
736
-
737
- for seq, score in beams:
738
- if seq[-1] == eos_id:
739
- all_candidates.append((seq, score))
740
- continue
741
-
742
- with torch.no_grad():
743
- input_seq = seq.unsqueeze(0)
744
- outputs = model.forward_inference(enc_output, input_seq)
745
- logits = outputs['logits'][:, -1, :]
746
-
747
- log_probs = F.log_softmax(logits, dim=-1)
748
- top_log_probs, top_ids = torch.topk(log_probs, beam_size, dim=-1)
749
-
750
- for i in range(beam_size):
751
- next_id = top_ids[0, i]
752
- log_prob = top_log_probs[0, i].item()
753
- new_seq = torch.cat([seq, next_id.view(1)])
754
- new_score = score + log_prob
755
- all_candidates.append((new_seq, new_score))
756
-
757
- ordered = sorted(all_candidates, key=lambda x: x[1], reverse=True)
758
- beams = ordered[:beam_size]
759
-
760
- if all(b[0][-1] == eos_id for b in beams):
761
- break
762
-
763
- best_beam = sorted(beams, key=lambda x: x[1] / len(x[0]), reverse=True)[0]
764
- best_seq = best_beam[0]
765
-
766
- return best_seq
767
-
768
-
769
- def predict_with_beam_search(model, image_tensor, vocab, beam_size, device):
770
- """
771
- ЗАПУСКАЕТ BEAM SEARCH НА ОСНОВЕ ГОТОВОГО ТЕНЗОРА ИЗОБРАЖЕНИЯ.
772
- """
773
- # <<< МЫ УДАЛИЛИ ВСЮ ЛОГИКУ ЗАГРУЗКИ И ТРАНСФОРМАЦИИ ИЗОБРАЖЕНИЯ ОТСЮДА >>>
774
- # Теперь функция просто принимает image_tensor
775
-
776
- # 2. Получение ID спецтокенов
777
- sos_id = vocab.stoi[SOS_TOKEN]
778
- eos_id = vocab.stoi[EOS_TOKEN]
779
- pad_id = vocab.stoi[PAD_TOKEN]
780
-
781
- # 3. Генерация (теперь передаем тензор напрямую)
782
- # Убедимся, что тензор на правильном устройстве
783
- predicted_ids = generate_beam_search(
784
- model, image_tensor.to(device), beam_size, MAX_SEQ_LEN, sos_id, eos_id, pad_id, device
785
- )
786
-
787
- # 4. Декодирование
788
- predicted_ids = predicted_ids.cpu().numpy()
789
- try:
790
- # Убираем все после EOS токена
791
- eos_index = list(predicted_ids).index(eos_id)
792
- predicted_ids = predicted_ids[:eos_index]
793
- except ValueError:
794
- pass # EOS не был сгенерирован
795
-
796
- # Декодируем, пропуская SOS токен в начале
797
- latex_string = vocab.decode(predicted_ids)
798
- return latex_string
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  import os
2
  import re
3
  import math
 
4
  import random
5
+ from collections import defaultdict
6
+ from einops import rearrange
7
  import numpy as np
8
+ import cv2
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
15
  from tqdm import tqdm
16
+ import torchvision.transforms.functional as TF
17
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
18
  TOKEN_REGEX = re.compile(
19
+ r"(\\[a-zA-Z]+(?:\*)?)|(\\\{|\\\}|\\.)|([0-9]+(?:\.[0-9]+)?(?:[eE][+-]?[0-9]+)?)|([A-Za-z]+)|([+\-*/=<>!~^_&|%])|([{}()\[\],.;:?'])|(\s+)|(\S)"
 
 
 
 
 
 
 
20
  )
21
 
22
+ def tokenize_latex(s: str) -> list[str]:
23
+ tokens_grouped = TOKEN_REGEX.findall(s)
 
 
 
 
 
 
 
24
  tokens = []
25
+ for group in tokens_grouped:
26
+ non_empty_token = next(filter(None, group), '')
27
+ if not non_empty_token.isspace():
28
+ tokens.append(non_empty_token)
 
 
 
 
 
 
 
29
  return tokens
30
 
31
+ class Vocab:
32
+ def __init__(self, expressions=None, min_freq=1):
33
+ self.itos = [PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN]
34
+ if expressions:
35
+ freq = defaultdict(int)
36
+ for expr in expressions:
37
+ for t in tokenize_latex(expr):
38
+ freq[t] += 1
39
+ self.itos.extend([tok for tok, count in sorted(freq.items(), key=lambda x: -x[1]) if count >= min_freq])
40
+ self.stoi = {tok: i for i, tok in enumerate(self.itos)}
41
+ self.pad_id = self.stoi[PAD_TOKEN]
42
+ self.sos_id = self.stoi[SOS_TOKEN]
43
+ self.eos_id = self.stoi[EOS_TOKEN]
44
+ self.unk_id = self.stoi[UNK_TOKEN]
45
+ def encode(self, tokens: list[str]) -> list[int]:
46
+ return [self.stoi.get(t, self.unk_id) for t in tokens]
47
+ def decode(self, ids: list[int]) -> str:
48
+ tokens = [self.itos[i] for i in ids if i not in {self.pad_id, self.sos_id, self.eos_id}]
49
+ return "".join(tokens)
50
+
51
+ class PosVocab:
52
+ def __init__(self):
53
+ self.itos = ['<PAD>', '<SOS>', '<EOS>', 'M', 'L', 'R']
54
+ self.stoi = {s: i for i, s in enumerate(self.itos)}
55
+ self.pad_id = self.stoi['<PAD>']
56
+ self.sos_id = self.stoi['<SOS>']
57
+
58
+ # =================================================================================
59
+ # 2. POSITION FOREST (ИСПРАВЛЕНО)
60
+ # =================================================================================
61
  class PositionForestEncoder:
62
+ def __init__(self, pos_vocab, max_len=MAX_SEQ_LEN, max_identifier_len=MAX_IDENTIFIER_LEN):
63
+ self.pos_vocab = pos_vocab
64
+ self.max_len = max_len
65
+ self.max_identifier_len = max_identifier_len
66
+ self._cache = {}
67
+
68
+ def _get_pos_strings(self, tokens):
69
+ """
70
+ ИСПРАВЛЕНО: Полностью рекурсивный парсер, который точнее следует логике
71
+ построения дерева позиций для вложенных структур.
72
+ """
73
  cache_key = tuple(tokens)
74
+ if cache_key in self._cache:
75
+ return self._cache[cache_key]
76
+
 
 
77
  T = len(tokens)
78
+ ids = ['M'] * T
79
+
80
+ def find_matching_brace(start_index):
81
+ depth = 1
82
+ for i in range(start_index + 1, T):
83
+ if tokens[i] == '{':
84
+ depth += 1
85
+ elif tokens[i] == '}':
86
+ depth -= 1
87
+ if depth == 0:
88
+ return i
89
+ return T - 1
90
+
91
+ def parse_recursive(start, end, current_pos_prefix):
92
+ i = start
93
+ while i < end:
94
  tok = tokens[i]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
95
 
96
+ # Обработка команд с одним аргументом: ^, _, \sqrt
97
+ if tok in ('^', '_', '\\sqrt') and i + 1 < end:
98
+ label = 'L' if tok in ('^', '\\sqrt') else 'R'
99
+ if tokens[i+1] == '{':
100
+ brace_end = find_matching_brace(i + 1)
101
+ for k in range(i + 2, brace_end):
102
+ ids[k] = current_pos_prefix + label
103
+ parse_recursive(i + 2, brace_end, current_pos_prefix + label)
104
+ i = brace_end
105
+ else: # Аргумент - один токен
106
+ ids[i+1] = current_pos_prefix + label
107
+ i += 1
108
+
109
+ # Обработка команд с двумя аргументами: \frac
110
+ elif tok == '\\frac' and i + 1 < end and tokens[i+1] == '{':
111
+ num_end = find_matching_brace(i + 1)
112
+ if num_end + 1 < end and tokens[num_end + 1] == '{':
113
+ den_end = find_matching_brace(num_end + 1)
114
+ # Числитель
115
+ for k in range(i + 2, num_end):
116
+ ids[k] = current_pos_prefix + 'L'
117
+ parse_recursive(i + 2, num_end, current_pos_prefix + 'L')
118
+ # Знаменатель
119
+ for k in range(num_end + 2, den_end):
120
+ ids[k] = current_pos_prefix + 'R'
121
+ parse_recursive(num_end + 2, den_end, current_pos_prefix + 'R')
122
+ i = den_end
123
+ else:
124
+ i = num_end
125
 
126
+ # Обработка \left, \right (они не добавляют уровень, но влияют на разметку)
127
+ # В данной реализации мы их просто пропускаем, как и другие группирующие символы
128
+ # Более сложная логика могла бы их учитывать, но это выходит за рамки статьи
129
 
130
+ i += 1
 
 
 
 
 
 
 
131
 
132
+ parse_recursive(0, T, 'M')
 
 
 
133
 
134
+ # Заменяем префиксы обратно на полные строки
135
+ final_ids = []
136
+ for i in range(T):
137
+ if ids[i] == 'M':
138
+ final_ids.append('M')
139
+ else:
140
+ final_ids.append(ids[i])
141
+
142
+ self._cache[cache_key] = final_ids
143
+ return final_ids
144
+
145
+ def process_formula(self, latex_str: str):
146
+ # Эта часть остается без изменений
147
+ tokens = tokenize_latex(latex_str)
148
+ tokens_gt = [SOS_TOKEN] + tokens[:self.max_len - 2] + [EOS_TOKEN]
149
+ pos_strings_raw = self._get_pos_strings(tokens[:self.max_len - 2])
150
+ nested_depth_gt = [0] + [min(len(pid) - 1, NUM_NESTED_LEVELS) for pid in pos_strings_raw] + [0]
151
+ rel_pos_gt = [0] + [1 if pid.endswith('L') else 2 if pid.endswith('R') else 0 for pid in pos_strings_raw] + [0]
152
+ pos_identifiers_ids = []
153
+ sos_pos_ids = [self.pos_vocab.stoi[c] for c in ['<SOS>', 'M', '<EOS>']]
154
+ pos_identifiers_ids.append(sos_pos_ids)
155
+ for pos_str in pos_strings_raw:
156
+ ids = [self.pos_vocab.stoi.get(c, 0) for c in pos_str]
157
+ ids = [self.pos_vocab.sos_id] + ids + [self.pos_vocab.stoi['<EOS>']]
158
+ pos_identifiers_ids.append(ids)
159
+ pos_identifiers_ids.append(sos_pos_ids)
160
+ final_len = len(tokens_gt)
161
+ tokens_gt_padded = tokens_gt + [PAD_TOKEN] * (self.max_len - final_len)
162
+ nested_depth_gt_padded = nested_depth_gt + [0] * (self.max_len - final_len)
163
+ rel_pos_gt_padded = rel_pos_gt + [0] * (self.max_len - final_len)
164
+ for i in range(len(pos_identifiers_ids)):
165
+ seq = pos_identifiers_ids[i][:self.max_identifier_len]
166
+ pos_identifiers_ids[i] = seq + [self.pos_vocab.pad_id] * (self.max_identifier_len - len(seq))
167
+ empty_pos_id_seq = [self.pos_vocab.pad_id] * self.max_identifier_len
168
+ padded_pos_ids = pos_identifiers_ids + [empty_pos_id_seq] * (self.max_len - final_len)
169
+ return {"tokens_gt": tokens_gt_padded, "pos_matrix": torch.tensor(padded_pos_ids, dtype=torch.long),
170
+ "nested_gt": torch.tensor(nested_depth_gt_padded, dtype=torch.long), "rel_pos_gt": torch.tensor(rel_pos_gt_padded, dtype=torch.long)}
171
+
172
+ # =================================================================================
173
+ # 3. DATASET & PREPROCESSING (без изменений)
174
+ # =================================================================================
175
+ class ResizeWithPadding:
176
+ def __init__(self, target_height, max_target_width, padding_value=255):
177
+ self.target_height, self.max_target_width, self.padding_value = target_height, max_target_width, padding_value
178
+ def __call__(self, img):
179
+ w, h = img.size; new_w = int(w * (self.target_height / h)); new_w = min(new_w, self.max_target_width)
180
+ img_resized = img.resize((new_w, self.target_height), Image.LANCZOS)
181
+ new_img = Image.new(img.mode, (self.max_target_width, self.target_height), self.padding_value)
182
+ new_img.paste(img_resized, (0, 0))
183
+ mask = torch.ones((self.target_height, self.max_target_width), dtype=torch.bool)
184
+ mask[:, :new_w] = False
185
+ return new_img, mask
186
+
187
+ class ScaleToLimitRange:
188
+ def __init__(self, w_lo: int, w_hi: int, h_lo: int, h_hi: int) -> None:
189
+ assert w_lo <= w_hi and h_lo <= h_hi
190
+ self.w_lo = w_lo
191
+ self.w_hi = w_hi
192
+ self.h_lo = h_lo
193
+ self.h_hi = h_hi
194
+
195
+ def __call__(self, img: np.ndarray) -> np.ndarray:
196
+ h, w = img.shape[:2]
197
+ scale_r = min(self.h_hi / h, self.w_hi / w)
198
+ if scale_r < 1.0:
199
+ # Картинка слишком большая, сжимаем ее пропорционально
200
+ img = cv2.resize(
201
+ img, None, fx=scale_r, fy=scale_r, interpolation=cv2.INTER_LINEAR
202
+ )
203
+ return img
204
 
205
+ scale_r = max(self.h_lo / h, self.w_lo / w)
206
+ if scale_r > 1.0:
207
+ # Картинка слишком маленькая, увеличиваем ее пропорционально
208
+ img = cv2.resize(
209
+ img, None, fx=scale_r, fy=scale_r, interpolation=cv2.INTER_LINEAR
210
+ )
211
+ return img
212
 
213
+ # Если картинка уже в нужных рамках, ничего не делаем
214
+ return img
215
+
216
+ class ScaleAugmentation:
217
+ def __init__(self, lo: float, hi: float) -> None:
218
+ assert lo <= hi
219
+ self.lo = lo
220
+ self.hi = hi
221
 
222
+ def __call__(self, img: np.ndarray) -> np.ndarray:
223
+ k = np.random.uniform(self.lo, self.hi)
224
+ img = cv2.resize(img, None, fx=k, fy=k, interpolation=cv2.INTER_LINEAR)
225
+ return img
226
 
227
+ class CROHMEDataset(Dataset):
228
+ def __init__(self, base_dir, caption_file, vocab, pos_vocab, is_train=True):
229
+ self.vocab = vocab
230
+ self.pfe = PositionForestEncoder(pos_vocab)
231
+ self.items = []
232
+ self.is_train = is_train
233
 
234
+ img_dir = os.path.join(base_dir, 'img')
235
+ formulas_path = os.path.join(base_dir, caption_file)
236
+ with open(formulas_path, 'r', encoding='utf-8') as f:
237
+ for line in f:
238
+ fid, latex = line.strip().split('\t')
239
+ self.items.append((os.path.join(img_dir, f"{fid}.bmp"), latex))
240
 
241
+ # --- Инициализируем наши классы трансформаций ---
242
+ H_MIN, H_MAX = 32, 256
243
+ W_MIN, W_MAX = 32, 512
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
244
 
245
+ # Эти классы будут вызываться вручную
246
+ self.scale_augmenter = None
247
+ if self.is_train:
248
+ # Эта трансформация работает с NUMPY
249
+ self.scale_augmenter = ScaleAugmentation(0.8, 1.2)
250
 
251
+ # Эта трансформация тоже работает с NUMPY
252
+ self.size_controller = ScaleToLimitRange(h_lo=H_MIN, h_hi=H_MAX, w_lo=W_MIN, w_hi=W_MAX)
253
+
254
+ def __len__(self):
255
+ return len(self.items)
256
+
257
+ def __getitem__(self, idx):
258
+ path, latex = self.items[idx]
259
+ try:
260
+ # 1. ЧИТАЕМ КАРТИНКУ СРАЗУ В NUMPY МАССИВ. БОЛЬШЕ НИКАКИХ PIL.OPEN
261
+ img_np = cv2.imread(path, cv2.IMREAD_GRAYSCALE)
262
+ if img_np is None:
263
+ raise FileNotFoundError()
264
+ except Exception:
265
+ # Если файл битый, берем следующий
266
+ return self.__getitem__((idx + 1) % len(self))
267
+
268
+ # --- НАШ РУЧНОЙ ПАЙПЛАЙН ---
269
 
270
+ # Шаг А: Применяем ScaleAugmentation (numpy -> numpy)
271
+ if self.scale_augmenter:
272
+ img_np = self.scale_augmenter(img_np)
273
 
274
+ # Шаг Б: Применяем ScaleToLimitRange (numpy -> numpy)
275
+ # Он применяется всегда, и для train, и для val
276
+ img_np = self.size_controller(img_np)
 
 
 
 
 
277
 
278
+ # Шаг В: Конвертируем в PIL только в самом конце, чтобы отдать в collate_fn
279
+ final_img_pil = Image.fromarray(img_np)
280
+
281
+ # --- Конец пайплайна ---
282
 
283
+ # Остальной код без изменений
284
+ data = self.pfe.process_formula(latex)
285
+ token_ids = torch.tensor(self.vocab.encode(data["tokens_gt"]), dtype=torch.long)
286
+
287
+ return (final_img_pil, token_ids, data["pos_matrix"], data["nested_gt"], data["rel_pos_gt"], latex)
288
+ # =================================================================================
289
+ # 4. ENCODER (без изменений)
290
+ # =================================================================================
291
+ class ImgPosEnc(nn.Module):
292
+ def __init__(self, d_model: int, temperature: float = 10000.0, normalize: bool = True, scale: float = None):
293
  super().__init__()
294
+ if d_model % 4 != 0: raise ValueError(f"d_model ({d_model}) должен быть кратен 4.")
295
  self.d_model = d_model
296
+ self.temperature = temperature
297
+ self.normalize = normalize
298
+ if scale is None: scale = 2 * math.pi
299
+ self.scale = scale
300
+
301
+ def forward(self, x: torch.Tensor, mask: torch.BoolTensor) -> torch.Tensor:
302
+ not_mask = ~mask
303
+ y_embed = not_mask.cumsum(1, dtype=torch.float32)
304
+ x_embed = not_mask.cumsum(2, dtype=torch.float32)
305
+
306
+ if self.normalize:
307
+ eps = 1e-6
308
+ y_embed = (y_embed / (y_embed[:, -1:, :] + eps)) * self.scale
309
+ x_embed = (x_embed / (x_embed[:, :, -1:] + eps)) * self.scale
310
+
311
+ dim_t_half = self.d_model // 2
312
+ dim_t = torch.arange(dim_t_half, dtype=torch.float32, device=x.device)
313
+ dim_t = self.temperature ** (2 * (dim_t // 2) / dim_t_half)
314
+
315
+ pos_x = x_embed[:, :, :, None] / dim_t
316
+ pos_y = y_embed[:, :, :, None] / dim_t
317
+
318
+ pos_x = torch.stack((pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4).flatten(3)
319
+ pos_y = torch.stack((pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4).flatten(3)
320
+
321
+ pos = torch.cat((pos_y, pos_x), dim=3)
322
+ return x + pos
323
+
324
+ class _Bottleneck(nn.Module):
325
+ def __init__(self, n_channels: int, growth_rate: int, use_dropout: bool):
326
+ super(_Bottleneck, self).__init__()
327
+ interChannels = 4 * growth_rate
328
+ self.conv1 = nn.Conv2d(n_channels, interChannels, kernel_size=1, bias=False)
329
+ self.bn1 = nn.BatchNorm2d(interChannels)
330
+ self.conv2 = nn.Conv2d(interChannels, growth_rate, kernel_size=3, padding=1, bias=False)
331
+ self.bn2 = nn.BatchNorm2d(growth_rate)
332
+ self.use_dropout = use_dropout
333
+ self.dropout = nn.Dropout(p=0.2)
334
+ def forward(self, x):
335
+ out = F.relu(self.bn1(self.conv1(x)), inplace=True)
336
+ if self.use_dropout: out = self.dropout(out)
337
+ out = F.relu(self.bn2(self.conv2(out)), inplace=True)
338
+ if self.use_dropout: out = self.dropout(out)
339
+ out = torch.cat((x, out), 1)
340
+ return out
341
+ class _Transition(nn.Module):
342
+ def __init__(self, n_channels: int, n_out_channels: int, use_dropout: bool):
343
+ super(_Transition, self).__init__()
344
+ self.conv1 = nn.Conv2d(n_channels, n_out_channels, kernel_size=1, bias=False)
345
+ self.bn1 = nn.BatchNorm2d(n_out_channels)
346
+ self.use_dropout = use_dropout
347
+ self.dropout = nn.Dropout(p=0.2)
348
+ def forward(self, x):
349
+ out = F.relu(self.bn1(self.conv1(x)), inplace=True)
350
+ if self.use_dropout: out = self.dropout(out)
351
+ out = F.avg_pool2d(out, 2, ceil_mode=True)
352
+ return out
353
+ class DenseNet(nn.Module):
354
+ def __init__(self, growth_rate: int, num_layers: int, reduction: float = 0.5, bottleneck: bool = True, use_dropout: bool = True):
355
+ super(DenseNet, self).__init__()
356
+ n_dense_blocks = num_layers
357
+ n_channels = 2 * growth_rate
358
+ self.conv1 = nn.Conv2d(1, n_channels, kernel_size=7, padding=3, stride=2, bias=False)
359
+ self.norm1 = nn.BatchNorm2d(n_channels)
360
+ self.dense1 = self._make_dense(n_channels, growth_rate, n_dense_blocks, bottleneck, use_dropout)
361
+ n_channels += n_dense_blocks * growth_rate
362
+ n_out_channels = int(math.floor(n_channels * reduction))
363
+ self.trans1 = _Transition(n_channels, n_out_channels, use_dropout)
364
+ n_channels = n_out_channels
365
+ self.dense2 = self._make_dense(n_channels, growth_rate, n_dense_blocks, bottleneck, use_dropout)
366
+ n_channels += n_dense_blocks * growth_rate
367
+ n_out_channels = int(math.floor(n_channels * reduction))
368
+ self.trans2 = _Transition(n_channels, n_out_channels, use_dropout)
369
+ n_channels = n_out_channels
370
+ self.dense3 = self._make_dense(n_channels, growth_rate, n_dense_blocks, bottleneck, use_dropout)
371
+ self.out_channels = n_channels + n_dense_blocks * growth_rate
372
+ self.post_norm = nn.BatchNorm2d(self.out_channels)
373
+ @staticmethod
374
+ def _make_dense(n_channels, growth_rate, n_dense_blocks, bottleneck, use_dropout):
375
+ layers = []
376
+ for _ in range(int(n_dense_blocks)):
377
+ if bottleneck:
378
+ layers.append(_Bottleneck(n_channels, growth_rate, use_dropout))
379
+ n_channels += growth_rate
380
+ return nn.Sequential(*layers)
381
+ def forward(self, x, x_mask):
382
+ out = self.conv1(x)
383
+ out = self.norm1(out)
384
+ out_mask = x_mask[:, ::2, ::2]
385
+ out = F.relu(out, inplace=True)
386
+ out = F.max_pool2d(out, 2, ceil_mode=True)
387
+ out_mask = out_mask[:, ::2, ::2]
388
+ out = self.dense1(out)
389
+ out = self.trans1(out)
390
+ out_mask = out_mask[:, ::2, ::2]
391
+ out = self.dense2(out)
392
+ out = self.trans2(out)
393
+ out_mask = out_mask[:, ::2, ::2]
394
+ out = self.dense3(out)
395
+ out = self.post_norm(out)
396
+ return out, out_mask
397
+ class Encoder(nn.Module):
398
+ def __init__(self, d_model, growth_rate, num_layers, dropout):
399
+ super().__init__()
400
+ self.densenet = DenseNet(growth_rate, num_layers)
401
+ self.feature_proj = nn.Conv2d(self.densenet.out_channels, d_model, 1)
402
+ self.pos_enc_2d = ImgPosEnc(d_model, normalize=True)
403
+ self.norm = nn.LayerNorm(d_model)
404
+ self.dropout = nn.Dropout(p=dropout)
405
 
406
+ def forward(self, img, img_mask):
407
+ feature, mask = self.densenet(img, img_mask)
408
+ feature = self.feature_proj(feature)
409
+
410
+ feature_permuted = feature.permute(0, 2, 3, 1)
411
+ pos_encoded_feature = self.pos_enc_2d(feature_permuted, mask)
412
+
413
+ normed_feature = self.norm(pos_encoded_feature)
414
+ dropped_feature = self.dropout(normed_feature)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
415
 
416
+ return rearrange(dropped_feature, "b h w d -> b (h w) d"), rearrange(mask, "b h w -> b (h w)")
 
 
 
 
417
 
418
+ # =================================================================================
419
+ # 5. DECODER & IAC (ИСПРАВЛЕНО)
420
+ # =================================================================================
421
+ class StandardCrossAttention(nn.Module):
422
+ """Обертка для стандартного MHA, чтобы интерфейс был как у IAC."""
423
+ def __init__(self, d_model, num_heads, dropout):
424
+ super().__init__()
425
+ self.mha = nn.MultiheadAttention(d_model, num_heads, dropout=dropout, batch_first=True)
426
+ def forward(self, q, k, v, ids, pad_mask): # `ids` не используется, но нужен для совместимости
427
+ return self.mha(q, k, v, key_padding_mask=pad_mask)[0]
428
 
429
+ class IAC(nn.Module):
430
+ def __init__(self, d_model, num_heads, dropout, vocab):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
431
  super().__init__()
 
 
 
 
 
 
 
 
432
  self.d_model = d_model
433
+ self.num_heads = num_heads
434
+ self.head_dim = d_model // num_heads
435
 
436
+ self.wq, self.wk, self.wv, self.wo = [nn.Linear(d_model, d_model, bias=False) for _ in range(4)]
437
+ self.phi_conv = nn.Conv2d(num_heads, num_heads, kernel_size=3, padding=1, groups=num_heads)
438
+ self.phi_linear = nn.Linear(self.num_heads, self.head_dim)
439
+ self.dropout = nn.Dropout(dropout)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
440
 
441
+ # ИСПРАВЛЕНО: Расширенный список структурных токенов
442
+ structure_symbols = {
443
+ '^', '_', '{', '}', '\\frac', '\\sqrt', '\\left', '\\right',
444
+ '\\big', '\\Big', '\\bigg', '\\Bigg', # Размеры скобок
445
+ '&', '\\\\', # Элементы матриц и таблиц
446
+ # Служебные токены также считаются структурными
447
+ PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN
448
  }
449
+ s_ids = [vocab.stoi[s] for s in structure_symbols if s in vocab.stoi]
450
+ self.register_buffer('structure_ids', torch.tensor(s_ids, dtype=torch.long))
451
+
452
+ def _compute_phi(self, accum_attn):
453
+ B, H, Lq, Lk = accum_attn.shape
454
+ # Проверка, что Lk является идеальным квадратом, чтобы избежать ошибок с sqrt
455
+ s_dim_float = math.sqrt(Lk)
456
+ if s_dim_float != int(s_dim_float): return 0 # Не можем сформировать квадратную карту внимания
457
+ s_dim = int(s_dim_float)
458
+
459
+ phi_in = rearrange(accum_attn, 'b h lq (s1 s2) -> (b lq) h s1 s2', s1=s_dim, s2=s_dim)
460
+ conv_out = self.phi_conv(phi_in)
461
+ phi_permuted = rearrange(conv_out, '(b lq) h s1 s2 -> b lq (s1 s2) h', b=B)
462
+ phi_features = self.phi_linear(phi_permuted)
463
+ return phi_features.unsqueeze(1)
464
+
465
+ def forward(self, q, k, v, ids, pad_mask):
466
+ B, Lq, _ = q.shape; Lk = k.shape[1]
467
+ Q = self.wq(q).view(B, Lq, self.num_heads, self.head_dim).transpose(1, 2)
468
+ K = self.wk(k).view(B, Lk, self.num_heads, self.head_dim).transpose(1, 2)
469
+ V = self.wv(v).view(B, Lk, self.num_heads, self.head_dim).transpose(1, 2)
470
+ scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim)
471
+
472
+ temp_attn_weights = F.softmax(scores, dim=-1).detach()
473
+ indicator = (~(ids.unsqueeze(-1) == self.structure_ids).any(-1)).float().view(B, 1, Lq, 1)
474
+ masked_attn = temp_attn_weights * indicator
475
+ shifted = torch.zeros_like(masked_attn)
476
+ if Lq > 1: shifted[:, :, 1:, :] = masked_attn[:, :, :-1, :]
477
+ accumulated_attention = torch.cumsum(shifted, dim=2)
478
+
479
+ phi = self._compute_phi(accumulated_attention)
480
+ correction = (Q.unsqueeze(3) * phi).sum(dim=-1) if isinstance(phi, torch.Tensor) else 0
481
+ corrected_scores = scores - correction
482
+
483
+ if pad_mask is not None:
484
+ corrected_scores = corrected_scores.masked_fill(pad_mask.unsqueeze(1).unsqueeze(2), float('-inf'))
485
+ final_attn_weights = F.softmax(corrected_scores, dim=-1)
486
+ context = torch.matmul(self.dropout(final_attn_weights), V).transpose(1, 2).reshape(B, Lq, self.d_model)
487
+ return self.wo(context)
488
 
489
+ class EnhancedDecoderLayer(nn.Module):
490
+ def __init__(self, d_model, num_heads, d_ff, dropout, cross_attention_module):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
491
  super().__init__()
492
+ self.self_attn = nn.MultiheadAttention(d_model, num_heads, dropout=dropout, batch_first=True)
493
+ # ИСПРАВЛЕНО: Используем переданный модуль (Standard MHA или IAC)
494
+ self.cross_attn = cross_attention_module
495
+ self.ffn = nn.Sequential(nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model))
496
+ self.norm1, self.norm2, self.norm3 = [nn.LayerNorm(d_model) for _ in range(3)]
497
+ self.dropout = nn.Dropout(dropout)
498
 
499
+ def forward(self, x, enc_out, ids, tgt_mask, tgt_pad_mask, mem_pad_mask):
500
+ x = x + self.dropout(self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), attn_mask=tgt_mask, key_padding_mask=tgt_pad_mask)[0])
501
+ x = x + self.dropout(self.cross_attn(self.norm2(x), enc_out, enc_out, ids, mem_pad_mask))
502
+ x = x + self.dropout(self.ffn(self.norm3(x)))
503
+ return x
504
 
505
+ # =================================================================================
506
+ # 6. FULL POSFORMER MODEL (ИСПРАВЛЕНО)
507
+ # =================================================================================
508
+ class LabelSmoothingCrossEntropy(nn.Module):
509
+ """УЛУЧШЕНИЕ: Добавлено для борьбы с переобучением."""
510
+ def __init__(self, smoothing=0.0):
511
+ super(LabelSmoothingCrossEntropy, self).__init__()
512
+ self.smoothing = smoothing
513
+ def forward(self, x, target, ignore_index=-100):
514
+ confidence = 1. - self.smoothing
515
+ logprobs = F.log_softmax(x, dim=-1)
516
+ nll_loss = -logprobs.gather(dim=-1, index=target.unsqueeze(1)).squeeze(1)
517
+ smooth_loss = -logprobs.mean(dim=-1)
518
+ loss = confidence * nll_loss + self.smoothing * smooth_loss
519
+ mask = (target != ignore_index)
520
+ return (loss * mask.float()).sum() / mask.float().sum()
521
+
522
+ class PosFormer(nn.Module):
523
+ def __init__(self, main_vocab, pos_vocab):
524
+ super().__init__()
525
+ self.vocab, self.pos_vocab = main_vocab, pos_vocab
526
+ self.pad_id, self.sos_id, self.eos_id = main_vocab.pad_id, main_vocab.sos_id, main_vocab.eos_id
527
+ self.encoder = Encoder(D_MODEL, GROWTH_RATE, NUM_ENCODER_LAYERS, DROPOUT_RATE)
528
+ self.emb_tok = nn.Embedding(len(main_vocab.itos), D_MODEL)
529
+ self.pos_id_emb = nn.Embedding(len(pos_vocab.itos), D_MODEL)
530
+ self.xi_proj = nn.Sequential(nn.Linear(D_MODEL * MAX_IDENTIFIER_LEN, D_MODEL), nn.GELU(), nn.LayerNorm(D_MODEL))
531
+
532
+ # ИСПРАВЛЕНО: Создаем декодеры с разным типом внимания
533
+ decoder_layers = []
534
+ for i in range(NUM_DECODER_LAYERS):
535
+ use_iac = i > 0 # IAC на 2м и 3м слое (индексы 1 и 2)
536
+ cross_attn_module = IAC(D_MODEL, NUM_HEADS, DROPOUT_RATE, main_vocab) if use_iac \
537
+ else StandardCrossAttention(D_MODEL, NUM_HEADS, DROPOUT_RATE)
538
+ decoder_layers.append(
539
+ EnhancedDecoderLayer(D_MODEL, NUM_HEADS, D_FF, DROPOUT_RATE, cross_attn_module)
540
+ )
541
+ self.decoders = nn.ModuleList(decoder_layers)
542
 
543
+ self.dec_norm = nn.LayerNorm(D_MODEL)
544
+ self.head_tok = nn.Linear(D_MODEL, len(main_vocab.itos))
545
+ self.head_nested = nn.Linear(D_MODEL, NUM_NESTED_LEVELS + 1)
546
+ self.head_rel_pos = nn.Linear(D_MODEL, 3) # 0: M, 1: L, 2: R
547
 
548
+ # УЛУЧШЕНИЕ: Добавляем Label Smoothing в модель
549
+ self.loss_fn_rec = LabelSmoothingCrossEntropy(smoothing=0.05)
 
550
 
551
+ self._init_weights()
552
+ self.rel_pos_map = {0: 'M', 1: 'L', 2: 'R'}
553
 
554
+ def _init_weights(self):
555
+ for p in self.parameters():
556
+ if p.dim() > 1: nn.init.xavier_uniform_(p)
557
+
558
+ def forward(self, imgs, img_mask, pos_matrix, token_ids):
559
+ vis_feats, mem_pad_mask = self.encoder(imgs, img_mask)
560
+ pos_input = self.xi_proj(rearrange(self.pos_id_emb(pos_matrix[:, :-1, :]), 'b l d e -> b l (d e)'))
561
+ token_input_ids = token_ids[:, :-1]
562
+ token_embeds = self.emb_tok(token_input_ids)
563
+ tgt = pos_input + token_embeds
564
+ tgt_mask = torch.triu(torch.ones(tgt.size(1), tgt.size(1), device=tgt.device), 1).bool()
565
+ tgt_pad_mask = (token_input_ids == self.pad_id)
566
+ dec_out = self.dec_norm(self._decode(tgt, vis_feats, token_input_ids, tgt_mask, tgt_pad_mask, mem_pad_mask))
567
+ return self.head_tok(dec_out), self.head_nested(dec_out), self.head_rel_pos(dec_out)
568
+
569
+ def _decode(self, tgt, mem, ids, tgt_mask, tgt_pad_mask, mem_pad_mask):
570
+ for layer in self.decoders:
571
+ tgt = layer(tgt, mem, ids, tgt_mask, tgt_pad_mask, mem_pad_mask)
572
+ return tgt
573
+
574
+ def compute_loss(self, logits, targets):
575
+ token_logits, nested_logits, rel_pos_logits = logits
576
+ token_ids_gt, nested_gt, rel_pos_gt = targets
577
+ token_targets = token_ids_gt[:, 1:]
578
+ nested_targets = nested_gt[:, 1:]
579
+ rel_pos_targets = rel_pos_gt[:, 1:]
580
+
581
+ # УЛУЧШЕНИЕ: Используем Label Smoothing для рекогници-потерь
582
+ loss_rec = self.loss_fn_rec(
583
+ token_logits.reshape(-1, token_logits.size(-1)),
584
+ token_targets.reshape(-1),
585
+ ignore_index=self.pad_id
586
+ )
587
+
588
+ mask = (token_targets != self.pad_id).flatten()
589
+ if not mask.any():
590
+ return {'total': loss_rec, 'rec': loss_rec, 'pos': torch.tensor(0.0, device=token_logits.device)}
591
+
592
+ loss_nested = F.cross_entropy(nested_logits.reshape(-1, nested_logits.size(-1))[mask],
593
+ nested_targets.reshape(-1)[mask])
594
+ loss_rel = F.cross_entropy(rel_pos_logits.reshape(-1, rel_pos_logits.size(-1))[mask],
595
+ rel_pos_targets.reshape(-1)[mask])
596
+ loss_pos = loss_nested + loss_rel
597
+ total_loss = loss_rec + loss_pos
598
+
599
+ return {'total': total_loss, 'rec': loss_rec, 'pos': loss_pos}
600
+
601
+ # ... (Все методы генерации generate, generate_beam_search, _construct_next_pos_string остаются без изменений) ...
602
+ # Они будут работать лучше, так как модель обучается на более качественных данных
603
+ def _construct_next_pos_string(self, prev_pos_str: str, pred_nested_level: int, pred_rel_pos: int) -> str:
604
+ prev_level = len(prev_pos_str) - 1
605
+ new_pos_str = prev_pos_str[:pred_nested_level + 1]
606
+ if pred_nested_level > prev_level:
607
+ if len(new_pos_str) == prev_level + 1:
608
+ new_pos_str += self.rel_pos_map.get(pred_rel_pos, 'M')
609
+ return new_pos_str
610
+
611
+ @torch.no_grad()
612
+ def generate(self, imgs, max_gen_len=150):
613
+ self.eval()
614
+ B = imgs.shape[0]; device = imgs.device
615
+ img_mask = torch.zeros_like(imgs[:, 0, :, :], dtype=torch.bool)
616
+ vis_feats, mem_pad_mask = self.encoder(imgs, img_mask)
617
+ generated_ids = torch.full((B, 1), self.sos_id, dtype=torch.long, device=device)
618
+ pos_strings_T = [['M'] for _ in range(B)]
619
+ is_finished = torch.zeros(B, dtype=torch.bool, device=device)
620
+ for t in range(max_gen_len - 1):
621
+ token_embeds = self.emb_tok(generated_ids)
622
+ pos_ids_list = []
623
+ max_len_in_batch = max(len(p_list) for p_list in pos_strings_T)
624
+ for i in range(B):
625
+ batch_pos_ids = []
626
+ for p_str in pos_strings_T[i]:
627
+ ids = [self.pos_vocab.sos_id] + [self.pos_vocab.stoi.get(c,0) for c in p_str] + [self.pos_vocab.stoi['<EOS>']]
628
+ padded_ids = ids[:MAX_IDENTIFIER_LEN] + [self.pos_vocab.pad_id] * (MAX_IDENTIFIER_LEN - len(ids))
629
+ batch_pos_ids.append(torch.tensor(padded_ids, device=device))
630
+ while len(batch_pos_ids) < max_len_in_batch:
631
+ batch_pos_ids.append(torch.full((MAX_IDENTIFIER_LEN,), self.pos_vocab.pad_id, device=device))
632
+ pos_ids_list.append(torch.stack(batch_pos_ids))
633
+ pos_matrix = torch.stack(pos_ids_list)
634
+ pos_embeds = self.xi_proj(rearrange(self.pos_id_emb(pos_matrix), 'b l d e -> b l (d e)'))
635
+ tgt = token_embeds + pos_embeds
636
+ tgt_mask = torch.triu(torch.ones(tgt.size(1), tgt.size(1), device=device), 1).bool()
637
+ dec_out = self.dec_norm(self._decode(tgt, vis_feats, generated_ids, tgt_mask, None, mem_pad_mask))
638
+ last_step_out = dec_out[:, -1, :]
639
+ token_logits, nested_logits, rel_pos_logits = self.head_tok(last_step_out), self.head_nested(last_step_out), self.head_rel_pos(last_step_out)
640
+ next_token_id = torch.argmax(token_logits, dim=-1).unsqueeze(1)
641
+ pred_nested_level, pred_rel_pos = torch.argmax(nested_logits, dim=-1), torch.argmax(rel_pos_logits, dim=-1)
642
+ generated_ids = torch.cat([generated_ids, next_token_id], dim=1)
643
+ for i in range(B):
644
+ if not is_finished[i]:
645
+ new_pos_str = self._construct_next_pos_string(pos_strings_T[i][-1], pred_nested_level[i].item(), pred_rel_pos[i].item())
646
+ pos_strings_T[i].append(new_pos_str)
647
+ is_finished |= (next_token_id.squeeze(-1) == self.eos_id)
648
+ if is_finished.all(): break
649
+ return generated_ids
650
+
651
+ @torch.no_grad()
652
+ def generate_beam_search(self, imgs, beam_size=5, max_gen_len=100):
653
+
654
+ self.eval(); B = imgs.shape[0]; device = imgs.device
655
+ img_mask = torch.zeros_like(imgs[:, 0, :, :], dtype=torch.bool)
656
+ vis_feats, mem_pad_mask = self.encoder(imgs, img_mask)
657
+ vis_feats = vis_feats.repeat_interleave(beam_size, dim=0)
658
+ if mem_pad_mask is not None: mem_pad_mask = mem_pad_mask.repeat_interleave(beam_size, dim=0)
659
+ effective_batch_size = B * beam_size
660
+ generated_ids = torch.full((effective_batch_size, 1), self.sos_id, dtype=torch.long, device=device)
661
+ pos_strings_T = [['M'] for _ in range(effective_batch_size)]
662
+ log_scores = torch.zeros(effective_batch_size, device=device)
663
+ is_finished = torch.zeros(effective_batch_size, dtype=torch.bool, device=device)
664
+ for t in range(max_gen_len - 1):
665
+ if is_finished.all(): break
666
+ token_embeds = self.emb_tok(generated_ids)
667
+ pos_ids_list = []
668
+ for i in range(effective_batch_size):
669
+ batch_pos_ids = []
670
+ for p_str in pos_strings_T[i]:
671
+ ids = [self.pos_vocab.sos_id] + [self.pos_vocab.stoi.get(c, 0) for c in p_str] + [self.pos_vocab.stoi['<EOS>']]
672
+ padded_ids = ids[:MAX_IDENTIFIER_LEN] + [self.pos_vocab.pad_id] * (MAX_IDENTIFIER_LEN - len(ids))
673
+ batch_pos_ids.append(torch.tensor(padded_ids, device=device))
674
+ pos_ids_list.append(torch.stack(batch_pos_ids))
675
+ pos_matrix = torch.stack(pos_ids_list)
676
+ pos_embeds = self.xi_proj(rearrange(self.pos_id_emb(pos_matrix), 'b l d e -> b l (d e)'))
677
+ tgt = token_embeds + pos_embeds
678
+ tgt_mask = torch.triu(torch.ones(tgt.size(1), tgt.size(1), device=device), 1).bool()
679
+ dec_out = self.dec_norm(self._decode(tgt, vis_feats, generated_ids, tgt_mask, None, mem_pad_mask))
680
+ last_step_out = dec_out[:, -1, :]
681
+ token_logits, nested_logits, rel_pos_logits = self.head_tok(last_step_out), self.head_nested(last_step_out), self.head_rel_pos(last_step_out)
682
+ log_probs = F.log_softmax(token_logits, dim=-1)
683
+ if t > 0:
684
+ log_probs[is_finished] = -float('inf')
685
+ log_probs[is_finished, self.pad_id] = 0
686
+ total_scores = log_probs + log_scores.unsqueeze(1)
687
+ total_scores = total_scores.view(B, -1)
688
+ top_scores, top_indices = torch.topk(total_scores, beam_size, dim=1)
689
+ beam_indices = top_indices // len(self.vocab.itos)
690
+ token_indices = top_indices % len(self.vocab.itos)
691
+ batch_indices = torch.arange(B, device=device).view(-1, 1).repeat(1, beam_size)
692
+ beam_indices_abs = beam_indices + (batch_indices * beam_size)
693
+ generated_ids = generated_ids[beam_indices_abs.view(-1)]
694
+ pos_strings_T = [pos_strings_T[i] for i in beam_indices_abs.view(-1).tolist()]
695
+ generated_ids = torch.cat([generated_ids, token_indices.view(-1, 1)], dim=1)
696
+ pred_nested_levels = torch.argmax(nested_logits[beam_indices_abs.view(-1)], dim=-1)
697
+ pred_rel_poses = torch.argmax(rel_pos_logits[beam_indices_abs.view(-1)], dim=-1)
698
+ for i in range(effective_batch_size):
699
+ new_pos_str = self._construct_next_pos_string(pos_strings_T[i][-1], pred_nested_levels[i].item(), pred_rel_poses[i].item())
700
+ pos_strings_T[i].append(new_pos_str)
701
+ log_scores = top_scores.view(-1)
702
+ is_finished = is_finished[beam_indices_abs.view(-1)] | (token_indices.view(-1) == self.eos_id)
703
+ seq_lengths = (generated_ids != self.pad_id).sum(dim=1).float()
704
+ seq_lengths = torch.max(seq_lengths, torch.ones_like(seq_lengths))
705
+ normalized_scores = (log_scores / seq_lengths).view(B, beam_size)
706
+ best_beam_indices = torch.argmax(normalized_scores, dim=1)
707
+ final_indices = best_beam_indices + torch.arange(B, device=device) * beam_size
708
+ return generated_ids[final_indices]