Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
visualizer.py153 linesDownload Raw Back to util
1import numpy as np
2import os
3import ntpath #for handling Windows paths.
4import time
5from . import util # . = same package
6from . import html
7import visdom # for visualizing data.
8
9
10class Visualizer():
11    def __init__(self, opt):
12        self.display_id = opt.display_id #In Visdom, you can create multiple windows to display various types of information, such as images, graphs, or tables, during the training and testing of machine learning models. The display_id specifies which window (by ID) the visualizations should be displayed in.
13        self.use_html = opt.isTrain and not opt.no_html
14        #opt.isTrain==True means using train_options.py subclass
15        #HTML visualization in the context of deep learning involves generating web-based visual representations of the training process, model performance, and data. This can be particularly useful for monitoring, debugging, and sharing results. 
16        self.win_size = opt.display_winsize
17        self.name = opt.name
18        self.opt = opt
19        self.saved = False
20        if self.display_id > 0: #if diplay_id==0, disable Visdom
21            self.vis = visdom.Visdom(port=opt.display_port, ipv6=False)
22
23        if self.use_html: #logging training loss
24            self.web_dir = os.path.join(opt.checkpoints_dir, opt.name, 'web')
25            self.img_dir = os.path.join(self.web_dir, 'images')
26            print('create web directory %s...' % self.web_dir)
27            util.mkdirs([self.web_dir, self.img_dir])
28        self.log_name = os.path.join(opt.checkpoints_dir, opt.name, 'loss_log.txt')
29        with open(self.log_name, "a") as log_file:
30            now = time.strftime("%c")
31            log_file.write('================ Training Loss (%s) ================\n' % now)
32
33    def reset(self):
34        self.saved = False
35
36    # (mingcv)|visuals|: dictionary of images to display or save
37    def display_current_results(self, visuals, epoch, save_result):
38        if self.display_id > 0:  #(mingcv)show images in the browser
39            ncols = self.opt.display_single_pane_ncols
40            if ncols > 0:
41                h, w = next(iter(visuals.values())).shape[:2]
42                table_css = """<style>
43                        table {border-collapse: separate; border-spacing:4px; white-space:nowrap; text-align:center}
44                        table td {width: %dpx; height: %dpx; padding: 4px; outline: 4px solid black}
45                        </style>""" % (w, h)
46                title = self.name
47                label_html = ''
48                label_html_row = ''
49                nrows = int(np.ceil(len(visuals.items()) / ncols))
50                images = []
51                idx = 0
52                for label, image_numpy in visuals.items():
53                    label_html_row += '<td>%s</td>' % label
54                    images.append(image_numpy.transpose([2, 0, 1]))
55                    idx += 1
56                    if idx % ncols == 0:
57                        label_html += '<tr>%s</tr>' % label_html_row
58                        label_html_row = ''
59                white_image = np.ones_like(image_numpy.transpose([2, 0, 1])) * 255
60                while idx % ncols != 0:
61                    images.append(white_image)
62                    label_html_row += '<td></td>'
63                    idx += 1
64                if label_html_row != '':
65                    label_html += '<tr>%s</tr>' % label_html_row
66                # pane col = image row
67                self.vis.images(images, nrow=ncols, env=self.opt.name, win=self.display_id + 1,
68                                padding=2, opts=dict(title=title + ' images'))
69                label_html = '<table>%s</table>' % label_html
70                self.vis.text(table_css + label_html, env=self.opt.name, win=self.display_id + 2,
71                              opts=dict(title=title + ' labels'))
72            else:
73                idx = 1
74                for label, image_numpy in visuals.items():
75                    self.vis.image(image_numpy.transpose([2, 0, 1]), opts=dict(title=label),
76                                   env=self.opt.name,
77                                   win=self.display_id + idx)
78                    idx += 1
79
80        if self.use_html and (save_result or not self.saved):  # save images to a html file
81            self.saved = True
82            for label, image_numpy in visuals.items():
83                img_path = os.path.join(self.img_dir, 'epoch%.3d_%s.png' % (epoch, label))
84                util.save_image(image_numpy, img_path)
85            # update website
86            webpage = html.HTML(self.web_dir, 'Experiment name = %s' % self.name, reflesh=1)
87            for n in range(epoch, 0, -1):
88                webpage.add_header('epoch [%d]' % n)
89                ims = []
90                txts = []
91                links = []
92
93                for label, image_numpy in visuals.items():
94                    img_path = 'epoch%.3d_%s.png' % (n, label)
95                    ims.append(img_path)
96                    txts.append(label)
97                    links.append(img_path)
98                webpage.add_images(ims, txts, links, height=self.win_size)
99            webpage.save()
100
101    #(mingcv) errors: dictionary of error labels and values
102    def plot_current_errors(self, epoch, counter_ratio, opt, errors):
103        if not hasattr(self, 'plot_data'):
104            self.plot_data = {'X': [], 'Y': [], 'legend': list(errors.keys())}
105        self.plot_data['X'].append(epoch + counter_ratio)
106        self.plot_data['Y'].append([errors[k] for k in self.plot_data['legend']])
107        self.vis.line(
108            X=np.stack([np.array(self.plot_data['X'])] * len(self.plot_data['legend']), 1),
109            Y=np.array(self.plot_data['Y']),
110            opts={
111                'title': self.name + ' loss over time',
112                'legend': self.plot_data['legend'],
113                'xlabel': 'epoch',
114                'ylabel': 'loss'},
115            win=self.display_id,
116            env=self.opt.name)
117
118    # errors: same format as |errors| of plotCurrentErrors
119    def print_current_errors(self, epoch, i, errors, t):
120        message = '(epoch: %d, iters: %d, time: %.3f) ' % (epoch, i, t)
121        for k, v in errors.items():
122            message += '%s: %.3f ' % (k, v)
123
124        print(message)
125        with open(self.log_name, "a") as log_file:
126            log_file.write('%s\n' % message)
127
128    # (mingcv) save image to the disk
129    def save_images(self, webpage, visuals, image_path, aspect_ratio=1.0):
130        image_dir = webpage.get_image_dir()
131        short_path = ntpath.basename(image_path[0])
132        name = os.path.splitext(short_path)[0]
133
134        webpage.add_header(name)
135        ims = []
136        txts = []
137        links = []
138
139        for label, im in visuals.items():
140            image_name = '%s_%s.png' % (name, label)
141            save_path = os.path.join(image_dir, image_name)
142            h, w, _ = im.shape
143            if aspect_ratio > 1.0:
144                im = np.array(Image.fromarray(im).resize((h, int(w * aspect_ratio))))
145            if aspect_ratio < 1.0:
146                im = np.array(Image.fromarray(im).resize((h, int(h / aspect_ratio))))
147            util.save_image(im, save_path)
148
149            ims.append(image_name)
150            txts.append(label)
151            links.append(image_name)
152        webpage.add_images(ims, txts, links, height=self.win_size)
153