MCPcopy Create free account
hub / github.com/chibohe/text_recognition_toolbox / Recoginizer

Class Recoginizer

predict.py:26–92  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class Recoginizer(object):
27 def __init__(self):
28 self.config = build_config()
29 model = build_model(self.config)
30 device, gpu_count = build_device(self.config)
31 optimizer = build_optimizer(self.config, model)
32 model, optimizer, global_state = build_pretrained_weights(self.config, model, optimizer)
33 self.device = device
34 if self.config.Global.loss_type == 'ctc':
35 self.converter = CTCLabelConverter(self.config)
36 else:
37 self.converter = AttnLabelConverter(self.config)
38 self.model = model.to(self.device)
39
40 self.keep_ratio_with_pad = self.config.TrainReader.padding
41 self.channel = self.config.Global.image_shape[0]
42 self.imgH = self.config.Global.image_shape[1]
43 self.imgW = self.config.Global.image_shape[2]
44 self.num_steps = self.config.Global.batch_max_length + 1
45
46 def preprocess(self, image):
47 self.transform = transforms.ToTensor()
48
49 if self.keep_ratio_with_pad:
50 w, h = image.size
51 ratio = w / float(h)
52 if math.ceil(ratio * self.imgH) > self.imgW:
53 resized_image = image.resize((self.imgW, self.imgH), Image.BICUBIC)
54 resized_image = self.transform(resized_image)
55 imgP = resized_image.sub(0.5).div(0.5)
56 else:
57 resized_W = math.ceil(ratio * self.imgH)
58 resized_image = image.resize((resized_W, self.imgH), Image.BICUBIC)
59 resized_image = self.transform(resized_image)
60 resized_image = resized_image.sub(0.5).div(0.5)
61
62 c, h, w = resized_image.size()
63 imgP = torch.FloatTensor(*(self.channel, self.imgH, self.imgW)).fill_(0)
64 imgP[:, :, :w] = resized_image
65 imgP[:, :, w:] = resized_image[:, :, w - 1].unsqueeze(2).expand(c, h, self.imgW - w)
66 else:
67 resized_image = image.resize((self.imgW, self.imgH), Image.BICUBIC)
68 resized_image = self.transform(resized_image)
69 imgP = resized_image.sub(0.5).div(0.5)
70
71 imgP = imgP.unsqueeze(0)
72 return imgP
73
74 def predict(self, image_tensor):
75 self.model.eval()
76 with torch.no_grad():
77 image_tensor = image_tensor.to(self.device)
78 if self.config.Global.loss_type == 'ctc':
79 outputs = self.model(image_tensor)
80 else:
81 pseudo_text = None
82 outputs = self.model(image_tensor, pseudo_text)
83 outputs = outputs.softmax(dim=2).detach().cpu().numpy()

Callers 1

predict.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected