G-Madhuri commited on
Commit
a505446
Β·
verified Β·
1 Parent(s): 3736caa

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +23 -33
app.py CHANGED
@@ -27,7 +27,7 @@ logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(
27
  logger = logging.getLogger(__name__)
28
 
29
  # =========================
30
- # Setup PARSeq path - USE LOCAL FOLDER, NOT TORCH HUB
31
  # =========================
32
  current_dir = os.path.dirname(os.path.abspath(__file__))
33
  parseq_local_path = os.path.join(current_dir, 'parseq')
@@ -46,7 +46,6 @@ try:
46
  logger.info("βœ… Successfully imported Tokenizer and PARSeq from local folder")
47
  except ImportError as e:
48
  logger.error(f"Failed to import: {e}")
49
- logger.info(f"Contents of parseq folder: {os.listdir(parseq_local_path)}")
50
  raise
51
 
52
  warnings.filterwarnings('ignore')
@@ -81,20 +80,6 @@ transform = T.Compose([
81
  T.Normalize(mean=[0.5], std=[0.5])
82
  ])
83
 
84
- # =========================
85
- # Decode
86
- # =========================
87
- def decode_prediction(logits, tokenizer):
88
- pred_ids = logits.argmax(-1)[0]
89
- chars = []
90
- for t in pred_ids:
91
- t = t.item()
92
- if t == tokenizer.eos_id:
93
- break
94
- if t not in [tokenizer.pad_id, tokenizer.bos_id] and t < len(tokenizer._itos):
95
- chars.append(tokenizer._itos[t])
96
- return "".join(chars)
97
-
98
  # =========================
99
  # Model Cache
100
  # =========================
@@ -140,7 +125,7 @@ def load_model(model_path, lang_name):
140
  k = k.replace('module.', '')
141
  new_state_dict[k] = v
142
 
143
- # Create model with required parameters
144
  model = PARSeq(
145
  num_tokens=len(charset_str),
146
  max_label_length=100,
@@ -159,7 +144,12 @@ def load_model(model_path, lang_name):
159
  )
160
 
161
  # Load weights
162
- model.load_state_dict(new_state_dict, strict=False)
 
 
 
 
 
163
  model.tokenizer = tokenizer
164
  model = model.to(device)
165
  model.eval()
@@ -175,7 +165,7 @@ def load_model(model_path, lang_name):
175
  return None, None, None
176
 
177
  # =========================
178
- # Inference - FIXED FORWARD METHOD
179
  # =========================
180
  def inference_image(model, image, device, tokenizer):
181
  if image is None:
@@ -187,21 +177,21 @@ def inference_image(model, image, device, tokenizer):
187
  img_tensor = transform(image).unsqueeze(0).to(device)
188
 
189
  with torch.no_grad():
190
- # Call forward with both images and tokenizer
191
- logits = model(images=img_tensor, tokenizer=tokenizer)
 
192
 
193
- # Get predicted text from logits
194
- # For PARSeq, the output might be logits or the model might return predictions directly
195
- if isinstance(logits, tuple):
196
- logits = logits[0] # Sometimes returns (logits, attention_weights)
 
 
 
 
 
197
 
198
- predicted_text = decode_prediction(logits, tokenizer)
199
-
200
- probs = torch.softmax(logits, dim=-1)
201
- max_probs = probs.max(dim=-1)[0][0]
202
- avg_conf = max_probs[:len(predicted_text)].mean().item() if len(predicted_text) > 0 else 0
203
-
204
- return predicted_text, avg_conf
205
 
206
  # =========================
207
  # Get samples for specific language
@@ -299,7 +289,7 @@ def create_language_tab(language):
299
 
300
  text, conf = inference_image(model, image, device, tokenizer)
301
 
302
- if text == "":
303
  return "πŸ” No text detected in the image", ""
304
 
305
  return text, f"βœ… Confidence: {conf:.2%}"
 
27
  logger = logging.getLogger(__name__)
28
 
29
  # =========================
30
+ # Setup PARSeq path
31
  # =========================
32
  current_dir = os.path.dirname(os.path.abspath(__file__))
33
  parseq_local_path = os.path.join(current_dir, 'parseq')
 
46
  logger.info("βœ… Successfully imported Tokenizer and PARSeq from local folder")
47
  except ImportError as e:
48
  logger.error(f"Failed to import: {e}")
 
49
  raise
50
 
51
  warnings.filterwarnings('ignore')
 
80
  T.Normalize(mean=[0.5], std=[0.5])
81
  ])
82
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
83
  # =========================
84
  # Model Cache
85
  # =========================
 
125
  k = k.replace('module.', '')
126
  new_state_dict[k] = v
127
 
128
+ # Create model
129
  model = PARSeq(
130
  num_tokens=len(charset_str),
131
  max_label_length=100,
 
144
  )
145
 
146
  # Load weights
147
+ missing, unexpected = model.load_state_dict(new_state_dict, strict=False)
148
+ if missing:
149
+ logger.warning(f"Missing keys: {len(missing)}")
150
+ if unexpected:
151
+ logger.warning(f"Unexpected keys: {len(unexpected)}")
152
+
153
  model.tokenizer = tokenizer
154
  model = model.to(device)
155
  model.eval()
 
165
  return None, None, None
166
 
167
  # =========================
168
+ # Inference - Use model's generate method
169
  # =========================
170
  def inference_image(model, image, device, tokenizer):
171
  if image is None:
 
177
  img_tensor = transform(image).unsqueeze(0).to(device)
178
 
179
  with torch.no_grad():
180
+ # Use the model's generate method instead of forward
181
+ # This handles the tokenization internally
182
+ pred_str = model.generate(img_tensor, tokenizer)
183
 
184
+ # Calculate confidence (approximate)
185
+ try:
186
+ # Get logits for confidence estimation
187
+ logits = model(images=img_tensor, tokenizer=tokenizer)
188
+ probs = torch.softmax(logits, dim=-1)
189
+ max_probs = probs.max(dim=-1)[0][0]
190
+ avg_conf = max_probs[:len(pred_str[0])].mean().item() if len(pred_str[0]) > 0 else 0
191
+ except:
192
+ avg_conf = 0.5 # Default confidence if can't calculate
193
 
194
+ return pred_str[0] if isinstance(pred_str, list) else pred_str, avg_conf
 
 
 
 
 
 
195
 
196
  # =========================
197
  # Get samples for specific language
 
289
 
290
  text, conf = inference_image(model, image, device, tokenizer)
291
 
292
+ if text == "" or text is None:
293
  return "πŸ” No text detected in the image", ""
294
 
295
  return text, f"βœ… Confidence: {conf:.2%}"