Update app.py
Browse files
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
|
| 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
|
| 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 -
|
| 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 |
-
#
|
| 191 |
-
|
|
|
|
| 192 |
|
| 193 |
-
#
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
logits =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 197 |
|
| 198 |
-
|
| 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%}"
|