Commit 5e38e8c1 authored by Andoni Jimenez's avatar Andoni Jimenez
Browse files

Add clickable fits image widget. Refactor static image widget.

parent 0511a74a
Loading
Loading
Loading
Loading
+1 −0
Original line number Diff line number Diff line
@@ -3,3 +3,4 @@ pyqt5
pyqt5-tool
requests
Pillow
pyds9
+39 −23
Original line number Diff line number Diff line
@@ -3,7 +3,7 @@ from PyQt5 import QtWidgets, uic
from PyQt5.QtCore import Qt, QObject, QEvent, QSize, QCoreApplication
from PyQt5.QtGui import QPixmap, QIcon
import pandas as pd
from src import utils
from src import utils, widgets
import sys

import json
@@ -61,16 +61,20 @@ class Ui(QtWidgets.QMainWindow):
        policy = QtWidgets.QSizePolicy.Policy
        gb_img:QtWidgets.QWidget = self.findChild(QtWidgets.QGroupBox, 'gb_img')
        
        # Add clickable image widget
        widget = widgets.ClickableImage()
        gb_img.layout().addWidget(widget, 0, 0, 1, wcols/2) # Full row width
        wcount += wcols/2
        # Save widget:
        self.widget_clickimg = widget

        # Add image container widget
        widget = QtWidgets.QLabel(None)
        widget.setObjectName("imageLabel") 
        widget.setSizePolicy(QtWidgets.QSizePolicy(policy.Expanding, policy.Expanding))
        widget.setAlignment(Qt.AlignmentFlag.AlignHCenter | Qt.AlignmentFlag.AlignVCenter)
        gb_img.layout().addWidget(widget, 0, 0, 1, wcols) # Full row width
        wcount += wcols
        # Save image widget:
        widget = widgets.StaticImage()
        gb_img.layout().addWidget(widget, wcount//wcols, wcount%wcols, 1, wcols/2) # Full row width
        wcount += wcols/2
        # Save widget:
        self.defaultImgPath = Path("res") / Path("image_not_found.png")
        self.imageLabel = widget
        self.widget_staticimg = widget
        
        # Instances to widget saves
        self.rbg = {} # buttongroups
@@ -103,14 +107,14 @@ class Ui(QtWidgets.QMainWindow):
                elif wcols-(wcount%wcols) == 1: # last column
                    wpol = QtWidgets.QSizePolicy(policy.Maximum, policy.Preferred)
                else:
                    wpol = QtWidgets.QSizePolicy(policy.Minimum, policy.Maximum)
                    wpol = QtWidgets.QSizePolicy(policy.Minimum, policy.Preferred)
                
                gb = newGroupBox(id, name, wpol)
                # COMMENTBOX
                if tconf['type'] == 'text':
                    widget = QtWidgets.QPlainTextEdit(None)
                    widget.setObjectName(f"tb_{id}")
                    widget.setSizePolicy(QtWidgets.QSizePolicy(policy.Expanding, policy.Minimum))
                    widget.setSizePolicy(QtWidgets.QSizePolicy(policy.Expanding, policy.Preferred))
                    addTextBoxTexts(widget, tconf)
                    widget.installEventFilter(self) # Set enter as save
                    
@@ -207,6 +211,8 @@ class Ui(QtWidgets.QMainWindow):
        
        # Dynamic COLUMNS
        utils.IDS['FILE'] = list(df.columns)
        utils.IDS['WIDGET'] = ['fits_coords']
        
        utils.IDS['TB'] = list(self.tb.keys())
        for _, cb in self.cb.items():
            utils.IDS['CB'].extend(list(cb.keys()))
@@ -216,7 +222,7 @@ class Ui(QtWidgets.QMainWindow):
            utils.IDS['RB'].extend(list(rb.keys()))
        utils.IDS['RB'].extend([""])

        utils.COLUMNS = utils.IDS['FILE'] + utils.IDS['RBG'] + utils.IDS['CB'] + utils.IDS['TB']
        utils.COLUMNS = utils.IDS['FILE'] + utils.IDS['WIDGET'] + utils.IDS['RBG'] + utils.IDS['CB'] + utils.IDS['TB']
        for v in ['ra', 'dec']:
            try:
                utils.COLUMNS.remove(v)
@@ -265,6 +271,7 @@ class Ui(QtWidgets.QMainWindow):
        self.show()

        # self.load_row() # not needed first row loaded on fillList
        self.showImage(self.imgPath) # img appears streched, so refresh img

    # List helpers:

@@ -355,10 +362,10 @@ class Ui(QtWidgets.QMainWindow):
            if self.has_groups:
                grp = int(self.fileList.item(index, 0).text())
                gal = int(self.fileList.item(index, 1).text())
                item_index = (self.df['group']==grp) & (self.df['galaxy']==gal)
                item_index =  self.df.index[(self.df['group']==grp) & (self.df['galaxy']==gal)]
            else:
                gal = int(self.fileList.item(index, 0).text())
                item_index = self.df['galaxy']==gal
                item_index = self.df.index[self.df['galaxy']==gal]
                
            self.df.loc[item_index, 'processed'] = True
            if self.has_groups:
@@ -393,6 +400,9 @@ class Ui(QtWidgets.QMainWindow):
                self.df.loc[item_index, id] = tb.toPlainText().replace('\n', '')
            # self.df.loc[item_index, 'comment'] = self.commentBox.toPlainText().replace('\n', '')

            if hasattr(self, 'widget_clickimg'):
                self.df.at[item_index.item(), 'fits_coords'] = self.widget_clickimg.get_coords()  # not working with .loc 

            self.fileList.selectRow(index+1)

            utils.save_df(self.df)
@@ -465,6 +475,11 @@ class Ui(QtWidgets.QMainWindow):
            # else:
            #     self.commentBox.setPlainText('')
            
            if hasattr(self, 'widget_clickimg'):
                fname = item['fullpath'].item().parent / Path(item['fits'].item())
                coords = item['fits_coords'].item()
                self.widget_clickimg.new_file(fname, coords)

        except (AttributeError, KeyError) as e:
            print(f"WARNING:\tEmpty item. [{e}]")
            self.imgPath = self.defaultImgPath
@@ -517,15 +532,16 @@ class Ui(QtWidgets.QMainWindow):
    # Image helpers:

    def showImage(self, imageFile:str) -> None:
        
        if hasattr(self, 'widget_staticimg'):
            if imageFile and Path.is_file(imageFile):
            pixmap = QPixmap(str(imageFile))
                fname = str(imageFile)
            else:
            pixmap = QPixmap(str(self.defaultImgPath))

        w = self.imageLabel.width()
        h = self.imageLabel.height()
                fname = str(self.defaultImgPath)
            self.widget_staticimg.set_pixmap(QPixmap(fname))
            
        self.imageLabel.setPixmap(pixmap.scaled(w, h, Qt.KeepAspectRatio, Qt.SmoothTransformation))
        # if hasattr(self, 'widget_clickimg'):
        #     self.widget_clickimg.canvas.draw()

    def getIconCell(self, active:bool) -> QtWidgets.QWidget:
        iconLabel = QtWidgets.QLabel()
+12 −3
Original line number Diff line number Diff line
@@ -249,6 +249,10 @@ def expand_df(selectedFiles):
                                             'galaxy': int,
                                            }
                                )
        
        if 'fits_coords' in importData.columns:
            importData['fits_coords'] = importData['fits_coords'].fillna("[]").apply(lambda x: eval(x))
            
        checkColumnsMismatch(importData.columns.values)
        importData['processed'] = True
        importData['fullpath'] = ''
@@ -311,6 +315,10 @@ def newEntry(row:pd.Series) -> dict:
    for i, tbCol in enumerate(getTextBoxes()):
        entry.update({tbCol: ''})

    if 'fits' in row:
        entry.update({'fits': row.fits})
        entry.update({'fits_coords': []})
    
    entry.update(
        {
            'processed': False,
@@ -333,10 +341,11 @@ def save_df(df:pd.DataFrame) -> None:
    #                            df.loc[df['processed'] == True]])

    # Remove old values, keep last ones:
    
    s_cols = ['galaxy']
    if 'group' in processedItems.columns:
        exportData = processedItems.drop_duplicates(['group','galaxy'], keep='last').sort_values('group')
    else:
        exportData = processedItems.drop_duplicates(['galaxy'], keep='last').sort_values('galaxy')
        s_cols.insert(0,'group')
    exportData = processedItems.drop_duplicates(s_cols, keep='last').sort_values(by=s_cols)

    # Export final dataframe:
    exportData.to_csv(args.savefile, columns=getExportableColumns(),

src/widgets.py

0 → 100644
+159 −0
Original line number Diff line number Diff line
from PyQt5 import QtWidgets
from PyQt5.QtCore import Qt
import pandas as pd
import numpy as np
from src import utils

from matplotlib.backends.backend_qt5agg import (
    FigureCanvasQTAgg as FigureCanvas,
    NavigationToolbar2QT as NavigationToolbar
    )

from matplotlib.figure import Figure

from astropy.io import fits
from astropy.wcs import WCS
from astropy.coordinates import SkyCoord
from astropy import units as u
from astropy.visualization import ImageNormalize, BaseInterval, MinMaxInterval, PercentileInterval, ZScaleInterval, AsinhStretch, SqrtStretch, HistEqStretch

import pyds9 as ds9


policy = QtWidgets.QSizePolicy.Policy


class ClickableImage(QtWidgets.QWidget):
    
    def __init__(self, has_toolbar:bool = True, *args, **kwargs):
        super(ClickableImage, self).__init__(*args, **kwargs)
                
        self.setMinimumSize(150,170)
        self.setSizePolicy(QtWidgets.QSizePolicy(policy.Expanding, policy.Expanding))
        #self.setAlignment(Qt.AlignmentFlag.AlignHCenter | Qt.AlignmentFlag.AlignVCenter)

        # Create internal canvas
        layout = QtWidgets.QGridLayout(self)
        
        self.figure = Figure()
        self.canvas = FigureCanvas(self.figure)
        self.button = QtWidgets.QPushButton(None)
        self.button.setSizePolicy(QtWidgets.QSizePolicy(policy.Maximum, policy.Fixed))
        self.button.setText("Open DS9")
        self.button.clicked.connect(self.callback_ds9)
        
        if has_toolbar:
            self.toolbar = NavigationToolbar(self.canvas, coordinates=False)
            layout.addWidget(self.toolbar, 0, 0, 1, 1)
            layout.addWidget(self.button, 0, 1, 1, 1)
            
        layout.addWidget(self.canvas, 1, 0, 1, 2)        
        self.canvas.mpl_connect('button_press_event', self.callback_button_press) 

    # funtionality
    def new_file(self, file:str, coords:list = None) -> None:
        self.add_image(file)
        
        if coords is None or len(coords)==0:
            return
        hdu = self.projection
        for coord in coords:
            ra, dec = coord
            point = SkyCoord(ra, dec, unit=(u.deg, u.deg)).to_pixel(self.projection)
            self.add_point(point)
            
    def add_image(self, path_img:str):
        # add image
        try:
            with fits.open(path_img) as hdul:
                self.projection = WCS(hdul[0].header)
                data = hdul[0].data
                self.isfits = True
            self.path_img = path_img
        except OSError as e:
            print(e)
            return

        # Add main axes
        if hasattr(self, 'ax'):
            self.ax.remove()
        self.ax = self.figure.add_axes([0,0,1,1], projection=self.projection)
        self.ax.set_axis_off()
        
        # Add normalized image
        normalize = ImageNormalize(data, interval=PercentileInterval(99.5), stretch=AsinhStretch())
        self.ax.imshow(data, origin='lower', norm=normalize, aspect='equal')
        
        # Create point arrays and scatter plot
        self.coords = np.empty([0,2], dtype=float)
        self.points = np.empty([0,2], dtype=int)
        self._scatter = self.ax.scatter(None, None, color='r', marker='2')
        
        # Update canvas
        self.canvas.draw_idle()
    
    def update_points(self) -> None:
        # print(self.coords, self.points)
        self._scatter.set_offsets(self.points)
        self.canvas.draw_idle()
        
    def add_point(self, point) -> None:
        if self.isfits:
            hdu = self.projection
            coords = hdu.pixel_to_world(point[0], point[1])
            point_deg = np.round([coords.ra.deg, coords.dec.deg], 6)
            
        self.coords = np.insert(self.coords, 0, point_deg, axis=0)
        self.points = np.insert(self.points, 0, point, axis=0)
        self.update_points()
    
    def delete_point(self, point) -> None:
        if self.points.size == 0:
            return
        # find nearest point by euclidean distance
        idx_min = np.sum((self.points-point)**2, axis=1, keepdims=True).argmin(axis=0)
        
        self.coords = np.delete(self.coords, idx_min, axis=0)
        self.points = np.delete(self.points, idx_min, axis=0)
        self.update_points()
    
    def get_coords(self):
        """Convert coords list of lists to list of tuples and return"""
        return [tuple(c) for c in self.coords.tolist()]
    
    # button callbacks
    def callback_button_press(self, event) -> None:
        x, y = event.xdata, event.ydata
        if x is None or y is None: # out of bounds
            return
        point = [x, y]
        if event.dblclick and event.button == 1: # left doubleclick
            self.add_point(point)
        elif event.dblclick and event.button == 3: # right doubleclick
            self.delete_point(point)
    
    def callback_ds9(self):        
        d = ds9.DS9() # Open ds9 (this assumes no ds9 instance is yet running)
        d.set(f"file '{self.path_img}'") # Load file
        d.set('zoom to fit') # Zoom to fit
        # Change the colormap and scaling
        d.set('cmap bb')
        d.set('scale log')
        # Add a label
        #d.set('regions command {text 30 20 #text="Texto Ejemplo" font="times 18 bold"}')
      

class StaticImage(QtWidgets.QLabel):
    
    def __init__(self, *args, **kwargs):
        super(StaticImage, self).__init__(*args, **kwargs)
        
        self.setObjectName("imageLabel") 
        self.setMinimumSize(100,100)
        self.setSizePolicy(QtWidgets.QSizePolicy(policy.Expanding, policy.Expanding))
        self.setAlignment(Qt.AlignmentFlag.AlignHCenter | Qt.AlignmentFlag.AlignVCenter)
        
    def set_pixmap(self, pixmap) -> None:
        w = self.width()
        h = self.height()
        self.setPixmap(pixmap.scaled(w, h, Qt.KeepAspectRatio, Qt.SmoothTransformation))
 No newline at end of file