bugroup/Eye_Tracking_Drift_Correction
3
1import zipfile2import os3import plotly.express as px4import plotly.graph_objects as go5from torch.utils.data.dataloader import DataLoader as dl6import yaml7from io import StringIO8import torch as t9import numpy as np10import pandas as pd11from torch.utils.data import Dataset as torch_dset12from PIL import Image13import torchvision.transforms.functional as tvfunc14import json15from matplotlib import pyplot as plt16import matplotlib.patches as patches17from matplotlib.font_manager import FontProperties18import pathlib as pl19import matplotlib as mpl20import streamlit as st21from streamlit.runtime.uploaded_file_manager import UploadedFile22import einops as eo23import copy24 25# import stqdm26from tqdm.auto import tqdm27import time28import requests29 30from matplotlib.patches import Rectangle31from matplotlib import font_manager32from models import LitModel, EnsembleModel33from loss_functions import corn_label_from_logits34import classic_correction_algos as calgo35import analysis_funcs as anf36 37TEMP_FOLDER = pl.Path("results")38AVAILABLE_FONTS = [x.name for x in font_manager.fontManager.ttflist]39PLOTS_FOLDER = pl.Path("plots")40TEMP_FIGURE_STIMULUS_PATH = PLOTS_FOLDER / "temp_matplotlib_plot_stimulus.png"41all_fonts = [x.name for x in font_manager.fontManager.ttflist]42mpl.use("agg")43 44DIST_MODELS_FOLDER = pl.Path("models")45IMAGENET_MEAN = [0.485, 0.456, 0.406]46IMAGENET_STD = [0.229, 0.224, 0.225]47gradio_plots = pl.Path("plots")48 49event_strs = [50 "EFIX",51 "EFIX R",52 "EFIX L",53 "SSACC",54 "ESACC",55 "SFIX",56 "MSG",57 "SBLINK",58 "EBLINK",59 "BUTTON",60 "INPUT",61 "END",62 "START",63 "DISPLAY ON",64]65names_dict = {66 "SSACC": {"Descr": "Start of Saccade", "Pattern": "SSACC <eye > <stime>"},67 "ESACC": {68 "Descr": "End of Saccade",69 "Pattern": "ESACC <eye > <stime> <etime > <dur> <sxp > <syp> <exp > <eyp> <ampl > <pv >",70 },71 "SFIX": {"Descr": "Start of Fixation", "Pattern": "SFIX <eye > <stime>"},72 "EFIX": {"Descr": "End of Fixation", "Pattern": "EFIX <eye > <stime> <etime > <dur> <axp > <ayp> <aps >"},73 "SBLINK": {"Descr": "Start of Blink", "Pattern": "SBLINK <eye > <stime>"},74 "EBLINK": {"Descr": "End of Blink", "Pattern": "EBLINK <eye > <stime> <etime > <dur>"},75 "DISPLAY ON": {"Descr": "Actual start of Trial", "Pattern": "DISPLAY ON"},76}77metadata_strs = ["DISPLAY COORDS", "GAZE_COORDS", "FRAMERATE"]78 79ALGO_CHOICES = st.session_state["ALGO_CHOICES"] = [80 "warp",81 "regress",82 "compare",83 "attach",84 "segment",85 "split",86 "stretch",87 "chain",88 "slice",89 "cluster",90 "merge",91 "Wisdom_of_Crowds",92 "DIST",93 "DIST-Ensemble",94 "Wisdom_of_Crowds_with_DIST",95 "Wisdom_of_Crowds_with_DIST_Ensemble",96]97COLORS = px.colors.qualitative.Alphabet98 99 100class NumpyEncoder(json.JSONEncoder):101 "From https://stackoverflow.com/questions/26646362/numpy-array-is-not-json-serializable"102 103 def default(self, obj):104 if isinstance(obj, np.ndarray):105 return obj.tolist()106 elif isinstance(obj, pl.Path) or isinstance(obj, UploadedFile):107 return str(obj)108 return json.JSONEncoder.default(self, obj)109 110 111class DSet(torch_dset):112 def __init__(113 self,114 in_sequence: t.Tensor,115 chars_center_coords_padded: t.Tensor,116 out_categories: t.Tensor,117 trialslist: list,118 padding_list: list = None,119 padding_at_end: bool = False,120 return_images_for_conv: bool = False,121 im_partial_string: str = "fixations_chars_channel_sep",122 input_im_shape=[224, 224],123 ) -> None:124 super().__init__()125 126 self.in_sequence = in_sequence127 self.chars_center_coords_padded = chars_center_coords_padded128 self.out_categories = out_categories129 self.padding_list = padding_list130 self.padding_at_end = padding_at_end131 self.trialslist = trialslist132 self.return_images_for_conv = return_images_for_conv133 self.input_im_shape = input_im_shape134 if return_images_for_conv:135 self.im_partial_string = im_partial_string136 self.plot_files = [137 str(x["plot_file"]).replace("fixations_words", im_partial_string) for x in self.trialslist138 ]139 140 def __getitem__(self, index):141 142 if self.return_images_for_conv:143 im = Image.open(self.plot_files[index])144 if [im.size[1], im.size[0]] != self.input_im_shape:145 im = tvfunc.resize(im, self.input_im_shape)146 im = tvfunc.normalize(tvfunc.to_tensor(im), IMAGENET_MEAN, IMAGENET_STD)147 if self.chars_center_coords_padded is not None:148 if self.padding_list is not None:149 attention_mask = t.ones(self.in_sequence[index].shape[:-1], dtype=t.long)150 if self.padding_at_end:151 if self.padding_list[index] > 0:152 attention_mask[-self.padding_list[index] :] = 0153 else:154 attention_mask[: self.padding_list[index]] = 0155 if self.return_images_for_conv:156 return (157 self.in_sequence[index],158 self.chars_center_coords_padded[index],159 im,160 attention_mask,161 self.out_categories[index],162 )163 return (164 self.in_sequence[index],165 self.chars_center_coords_padded[index],166 attention_mask,167 self.out_categories[index],168 )169 else:170 if self.return_images_for_conv:171 return (172 self.in_sequence[index],173 self.chars_center_coords_padded[index],174 im,175 self.out_categories[index],176 )177 else:178 return (self.in_sequence[index], self.chars_center_coords_padded[index], self.out_categories[index])179 180 if self.padding_list is not None:181 attention_mask = t.ones(self.in_sequence[index].shape[:-1], dtype=t.long)182 if self.padding_at_end:183 if self.padding_list[index] > 0:184 attention_mask[-self.padding_list[index] :] = 0185 else:186 attention_mask[: self.padding_list[index]] = 0187 if self.return_images_for_conv:188 return (self.in_sequence[index], im, attention_mask, self.out_categories[index])189 else:190 return (self.in_sequence[index], attention_mask, self.out_categories[index])191 if self.return_images_for_conv:192 return (self.in_sequence[index], im, self.out_categories[index])193 else:194 return (self.in_sequence[index], self.out_categories[index])195 196 def __len__(self):197 if isinstance(self.in_sequence, t.Tensor):198 return self.in_sequence.shape[0]199 else:200 return len(self.in_sequence)201 202 203def download_url(url, target_filename):204 r = requests.get(url)205 open(target_filename, "wb").write(r.content)206 return 0207 208 209def asc_to_trial_ids(asc_file, close_gap_between_words=True):210 if "logger" in st.session_state:211 st.session_state["logger"].debug("asc_to_trial_ids entered")212 asc_encoding = ["ISO-8859-15", "UTF-8"][0]213 trials_dict, lines = file_to_trials_and_lines(214 asc_file, asc_encoding, close_gap_between_words=close_gap_between_words215 )216 217 trials_by_ids = {trials_dict[idx]["trial_id"]: trials_dict[idx] for idx in trials_dict["paragraph_trials"]}218 if hasattr(asc_file, "name"):219 if "logger" in st.session_state:220 st.session_state["logger"].info(f"Found {len(trials_by_ids)} trials in {asc_file.name}.")221 return trials_by_ids, lines222 223 224def get_trials_list(asc_file=None, close_gap_between_words=True):225 if "logger" in st.session_state:226 st.session_state["logger"].debug("get_trials_list entered")227 228 if asc_file == None:229 if "single_asc_file" in st.session_state.keys() and st.session_state["single_asc_file"] is not None:230 asc_file = st.session_state["single_asc_file"]231 else:232 if "logger" in st.session_state:233 st.session_state["logger"].warning("Asc file is None")234 return None235 236 if hasattr(asc_file, "name"):237 if "logger" in st.session_state:238 st.session_state["logger"].info(f"get_trials_list entered with asc_file {asc_file.name}")239 240 trials_by_ids, lines = asc_to_trial_ids(asc_file, close_gap_between_words=close_gap_between_words)241 trial_keys = list(trials_by_ids.keys())242 243 return trial_keys, trials_by_ids, lines, asc_file244 245 246def save_trial_to_json(trial, savename):247 if "dffix" in trial:248 trial.pop("dffix")249 with open(savename, "w", encoding="utf-8") as f:250 json.dump(trial, f, ensure_ascii=False, indent=4, cls=NumpyEncoder)251 252 253def export_csv(dffix, trial):254 if isinstance(dffix, dict):255 dffix = dffix["value"]256 trial_id = trial["trial_id"]257 savename = TEMP_FOLDER.joinpath(pl.Path(trial["fname"]).stem)258 trial_name = f"{savename}_{trial_id}_trial_info.json"259 csv_name = f"{savename}_{trial_id}.csv"260 dffix.to_csv(csv_name)261 if "logger" in st.session_state:262 st.session_state["logger"].info(f"Saved processed data as {csv_name}")263 save_trial_to_json(trial, trial_name)264 if "logger" in st.session_state:265 st.session_state["logger"].info(f"Saved processed trial data as {trial_name}")266 267 return csv_name, trial_name268 269 270def get_all_classic_preds(dffix, trial, classic_algos_cfg):271 corrections = []272 for algo, classic_params in copy.deepcopy(classic_algos_cfg).items():273 dffix = calgo.apply_classic_algo(dffix, trial, algo, classic_params)274 corrections.append(np.asarray(dffix.loc[:, f"y_{algo}"]))275 return dffix, corrections276 277 278def apply_woc(dffix, trial, corrections, algo_choice):279 280 corrected_Y = calgo.wisdom_of_the_crowd(corrections)281 dffix.loc[:, f"y_{algo_choice}"] = corrected_Y282 dffix[f"y_{algo_choice}_correction"] = (dffix.loc[:, f"y_{algo_choice}"] - dffix.loc[:, "y"]).round(1)283 corrected_line_nums = [trial["y_char_unique"].index(y) for y in corrected_Y]284 dffix.loc[:, f"line_num_y_{algo_choice}"] = corrected_line_nums285 return dffix286 287 288def calc_xdiff_ydiff(line_xcoords_no_pad, line_ycoords_no_pad, line_heights, allow_multiple_values=False):289 x_diffs = np.unique(np.diff(line_xcoords_no_pad))290 if len(x_diffs) == 1:291 x_diff = x_diffs[0]292 elif not allow_multiple_values:293 x_diff = np.min(x_diffs)294 else:295 x_diff = x_diffs296 297 if np.unique(line_ycoords_no_pad).shape[0] == 1:298 return x_diff, line_heights[0]299 y_diffs = np.unique(np.diff(line_ycoords_no_pad))300 if len(y_diffs) == 1:301 y_diff = y_diffs[0]302 elif len(y_diffs) == 0:303 y_diff = 0304 elif not allow_multiple_values:305 y_diff = np.min(y_diffs)306 else:307 y_diff = y_diffs308 return x_diff, y_diff309 310 311def add_words(trial, close_gap_between_words=True):312 chars_list_reconstructed = []313 words_list = []314 word_start_idx = 0315 chars_df = pd.DataFrame(trial["chars_list"])316 chars_df["char_width"] = chars_df.char_xmax - chars_df.char_xmin317 space_width = chars_df.loc[chars_df["char"] == " ", "char_width"].mean()318 319 for idx, char_dict in enumerate(trial["chars_list"]):320 on_line_num = char_dict["assigned_line"]321 chars_list_reconstructed.append(char_dict)322 if (323 char_dict["char"] in [" ", ",", ";", ".", ":"]324 or (325 len(chars_list_reconstructed) > 2326 and (chars_list_reconstructed[-1]["char_xmin"] < chars_list_reconstructed[-2]["char_xmin"])327 )328 or len(chars_list_reconstructed) == len(trial["chars_list"])329 ):330 triggered = True331 word_xmin = chars_list_reconstructed[word_start_idx]["char_xmin"]332 word_xmax = chars_list_reconstructed[-2]["char_xmax"]333 word_ymin = chars_list_reconstructed[word_start_idx]["char_ymin"]334 word_ymax = chars_list_reconstructed[word_start_idx]["char_ymax"]335 word_x_center = (word_xmax - word_xmin) / 2 + word_xmin336 word_y_center = (word_ymax - word_ymin) / 2 + word_ymin337 word = "".join(338 [339 chars_list_reconstructed[idx]["char"]340 for idx in range(word_start_idx, len(chars_list_reconstructed) - 1)341 ]342 )343 assigned_line = chars_list_reconstructed[word_start_idx]["assigned_line"]344 345 word_dict = dict(346 word=word,347 word_xmin=word_xmin,348 word_xmax=word_xmax,349 word_ymin=word_ymin,350 word_ymax=word_ymax,351 word_x_center=word_x_center,352 word_y_center=word_y_center,353 assigned_line=assigned_line,354 )355 if char_dict["char"] != " ":356 word_start_idx = idx357 else:358 word_start_idx = idx + 1359 words_list.append(word_dict)360 else:361 triggered = False362 last_letter_in_word = word_dict["word"][-1]363 last_letter_in_chars_list_reconstructed = char_dict["char"]364 if last_letter_in_word != last_letter_in_chars_list_reconstructed:365 word_dict = dict(366 word=char_dict["char"],367 word_xmin=char_dict["char_xmin"],368 word_xmax=char_dict["char_xmax"],369 word_ymin=char_dict["char_ymin"],370 word_ymax=char_dict["char_ymax"],371 word_x_center=char_dict["char_x_center"],372 word_y_center=char_dict["char_y_center"],373 assigned_line=assigned_line,374 )375 words_list.append(word_dict)376 377 if close_gap_between_words:378 for widx in range(1, len(words_list)):379 if words_list[widx]["assigned_line"] == words_list[widx - 1]["assigned_line"]:380 word_sep_half_width = (words_list[widx]["word_xmin"] - words_list[widx - 1]["word_xmax"]) / 2381 words_list[widx - 1]["word_xmax"] = words_list[widx - 1]["word_xmax"] + word_sep_half_width382 words_list[widx]["word_xmin"] = words_list[widx]["word_xmin"] - word_sep_half_width383 384 return words_list385 386 387def asc_lines_to_trials_by_trail_id(388 lines: list, paragraph_trials_only=False, fname: str = "", close_gap_between_words=True389) -> dict:390 if hasattr(fname, "name"):391 fname = fname.name392 fps = -999393 display_coords = -999394 trials_dict = dict(paragraph_trials=[], paragraph_trial_IDs=[])395 trial_idx = -1396 removed_trial_ids = []397 for idx, l in enumerate(lines):398 parts = l.strip().split(" ")399 if "TRIALID" in l:400 trial_id = parts[-1]401 trial_idx += 1402 if trial_id[0] == "F":403 trial_is = "question"404 elif trial_id[0] == "P":405 trial_is = "practice"406 else:407 trial_is = "paragraph"408 trials_dict["paragraph_trials"].append(trial_idx)409 trials_dict["paragraph_trial_IDs"].append(trial_id)410 trials_dict[trial_idx] = dict(trial_id=trial_id, trial_id_idx=idx, trial_is=trial_is, filename=fname)411 last_trial_skipped = False412 413 elif "TRIAL_RESULT" in l or "stop_trial" in l:414 trials_dict[trial_idx]["trial_result_idx"] = idx415 trials_dict[trial_idx]["trial_result_timestamp"] = int(parts[0].split("\t")[1])416 if len(parts) > 2:417 trials_dict[trial_idx]["trial_result_number"] = int(parts[2])418 elif "DISPLAY COORDS" in l and isinstance(display_coords, int):419 display_coords = (float(parts[-4]), float(parts[-3]), float(parts[-2]), float(parts[-1]))420 elif "GAZE_COORDS" in l and isinstance(display_coords, int):421 display_coords = (float(parts[-4]), float(parts[-3]), float(parts[-2]), float(parts[-1]))422 elif "FRAMERATE" in l:423 l_idx = parts.index(metadata_strs[2])424 fps = float(parts[l_idx + 1])425 elif "TRIAL ABORTED" in l or "TRIAL REPEATED" in l:426 if not last_trial_skipped:427 if trial_is == "paragraph":428 trials_dict["paragraph_trials"].remove(trial_idx)429 trial_idx -= 1430 removed_trial_ids.append(trial_id)431 last_trial_skipped = True432 433 if paragraph_trials_only:434 trials_dict_temp = trials_dict.copy()435 for k in trials_dict_temp.keys():436 if k not in ["paragraph_trials"] + trials_dict_temp["paragraph_trials"]:437 trials_dict.pop(k)438 if len(trials_dict_temp["paragraph_trials"]):439 trial_idx = trials_dict_temp["paragraph_trials"][-1]440 else:441 return trials_dict442 trials_dict["display_coords"] = display_coords443 trials_dict["fps"] = fps444 trials_dict["max_trial_idx"] = trial_idx445 enum = trials_dict["paragraph_trials"] if "paragraph_trials" in trials_dict.keys() else range(len(trials_dict))446 for trial_idx in enum:447 if trial_idx not in trials_dict.keys():448 continue449 chars_list = []450 if "display_coords" not in trials_dict[trial_idx].keys():451 trials_dict[trial_idx]["display_coords"] = trials_dict["display_coords"]452 trial_start_idx = trials_dict[trial_idx]["trial_id_idx"]453 trial_end_idx = trials_dict[trial_idx]["trial_result_idx"]454 trial_lines = lines[trial_start_idx:trial_end_idx]455 for idx, l in enumerate(trial_lines):456 parts = l.strip().split(" ")457 if "START" in l and " MSG" not in l:458 trials_dict[trial_idx]["start_idx"] = trial_start_idx + idx + 7459 trials_dict[trial_idx]["start_time"] = int(parts[0].split("\t")[1])460 elif "END" in l and "ENDBUTTON" not in l and " MSG" not in l:461 trials_dict[trial_idx]["end_idx"] = trial_start_idx + idx - 2462 trials_dict[trial_idx]["end_time"] = int(parts[0].split("\t")[1])463 elif "SYNCTIME" in l:464 trials_dict[trial_idx]["synctime"] = trial_start_idx + idx465 trials_dict[trial_idx]["synctime_time"] = int(parts[0].split("\t")[1])466 elif "GAZE TARGET OFF" in l:467 trials_dict[trial_idx]["gaze_targ_off_time"] = int(parts[0].split("\t")[1])468 elif "GAZE TARGET ON" in l:469 trials_dict[trial_idx]["gaze_targ_on_time"] = int(parts[0].split("\t")[1])470 elif "DISPLAY_SENTENCE" in l: # some .asc files seem to use this471 trials_dict[trial_idx]["gaze_targ_on_time"] = int(parts[0].split("\t")[1])472 elif "REGION CHAR" in l:473 rg_idx = parts.index("CHAR")474 if len(parts[rg_idx:]) > 8:475 char = " "476 idx_correction = 1477 elif len(parts[rg_idx:]) == 3:478 char = " "479 if "REGION CHAR" not in trial_lines[idx + 1]:480 parts = trial_lines[idx + 1].strip().split(" ")481 idx_correction = -rg_idx - 4482 else:483 char = parts[rg_idx + 3]484 idx_correction = 0485 try:486 char_dict = {487 "char": char,488 "char_xmin": float(parts[rg_idx + 4 + idx_correction]),489 "char_ymin": float(parts[rg_idx + 5 + idx_correction]),490 "char_xmax": float(parts[rg_idx + 6 + idx_correction]),491 "char_ymax": float(parts[rg_idx + 7 + idx_correction]),492 }493 char_dict["char_y_center"] = (char_dict["char_ymax"] - char_dict["char_ymin"]) / 2 + char_dict[494 "char_ymin"495 ]496 char_dict["char_x_center"] = (char_dict["char_xmax"] - char_dict["char_xmin"]) / 2 + char_dict[497 "char_xmin"498 ]499 chars_list.append(char_dict)500 except Exception as e:501 if "logger" in st.session_state:502 st.session_state["logger"].warning(f"char_dict creation failed for parts {parts}")503 if "logger" in st.session_state:504 st.session_state["logger"].warning(e)505 506 if "gaze_targ_on_time" in trials_dict[trial_idx]:507 trials_dict[trial_idx]["trial_start_time"] = trials_dict[trial_idx]["gaze_targ_on_time"]508 else:509 trials_dict[trial_idx]["trial_start_time"] = trials_dict[trial_idx]["start_time"]510 511 if len(chars_list) > 0:512 line_ycoords = []513 for idx in range(len(chars_list)):514 chars_list[idx]["char_line_y"] = (515 chars_list[idx]["char_ymax"] - chars_list[idx]["char_ymin"]516 ) / 2 + chars_list[idx]["char_ymin"]517 if chars_list[idx]["char_line_y"] not in line_ycoords:518 line_ycoords.append(chars_list[idx]["char_line_y"])519 for idx in range(len(chars_list)):520 chars_list[idx]["assigned_line"] = line_ycoords.index(chars_list[idx]["char_line_y"])521 522 line_heights = [x["char_ymax"] - x["char_ymin"] for x in chars_list]523 line_xcoords_all = [x["char_x_center"] for x in chars_list]524 line_xcoords_no_pad = np.unique(line_xcoords_all)525 526 line_ycoords_all = [x["char_y_center"] for x in chars_list]527 line_ycoords_no_pad = np.unique(line_ycoords_all)528 529 trials_dict[trial_idx]["x_char_unique"] = list(line_xcoords_no_pad)530 trials_dict[trial_idx]["y_char_unique"] = list(line_ycoords_no_pad)531 x_diff, y_diff = calc_xdiff_ydiff(532 line_xcoords_no_pad, line_ycoords_no_pad, line_heights, allow_multiple_values=False533 )534 trials_dict[trial_idx]["x_diff"] = float(x_diff)535 trials_dict[trial_idx]["y_diff"] = float(y_diff)536 trials_dict[trial_idx]["num_char_lines"] = len(line_ycoords_no_pad)537 trials_dict[trial_idx]["line_heights"] = line_heights538 trials_dict[trial_idx]["chars_list"] = chars_list539 540 words_list = add_words(trials_dict[trial_idx], close_gap_between_words=close_gap_between_words)541 trials_dict[trial_idx]["words_list"] = words_list542 543 return trials_dict544 545 546def file_to_trials_and_lines(uploaded_file, asc_encoding: str = "ISO-8859-15", close_gap_between_words=True):547 if isinstance(uploaded_file, str) or isinstance(uploaded_file, pl.Path):548 with open(uploaded_file, "r", encoding=asc_encoding) as f:549 lines = f.readlines()550 else:551 stringio = StringIO(uploaded_file.getvalue().decode(asc_encoding))552 loaded_str = stringio.read()553 lines = loaded_str.split("\n")554 trials_dict = asc_lines_to_trials_by_trail_id(555 lines, True, uploaded_file, close_gap_between_words=close_gap_between_words556 )557 558 if "paragraph_trials" not in trials_dict.keys() and "trial_is" in trials_dict[0].keys():559 paragraph_trials = []560 for k in range(trials_dict["max_trial_idx"]):561 if trials_dict[k]["trial_is"] == "paragraph":562 paragraph_trials.append(k)563 trials_dict["paragraph_trials"] = paragraph_trials564 565 enum = (566 trials_dict["paragraph_trials"]567 if "paragraph_trials" in trials_dict.keys()568 else range(trials_dict["max_trial_idx"])569 )570 for k in enum:571 if "chars_list" in trials_dict[k].keys():572 max_line = trials_dict[k]["chars_list"][-1]["assigned_line"]573 words_on_lines = {x: [] for x in range(max_line + 1)}574 [words_on_lines[x["assigned_line"]].append(x["char"]) for x in trials_dict[k]["chars_list"]]575 sentence_list = ["".join([s for s in v]) for idx, v in words_on_lines.items()]576 text = sentence_list[0] + "\n".join([x for x in sentence_list[1:]])577 trials_dict[k]["sentence_list"] = sentence_list578 trials_dict[k]["text"] = text579 trials_dict[k]["max_line"] = max_line580 581 return trials_dict, lines582 583 584def get_plot_props(trial, available_fonts):585 if "font" in trial.keys():586 font = trial["font"]587 font_size = trial["font_size"]588 if font not in available_fonts:589 font = "DejaVu Sans Mono"590 else:591 font = "DejaVu Sans Mono"592 font_size = 21593 dpi = 100594 if "display_coords" in trial.keys():595 screen_res = (trial["display_coords"][2], trial["display_coords"][3])596 else:597 screen_res = (1920, 1080)598 return font, font_size, dpi, screen_res599 600 601def trial_to_dfs(602 trial: dict, lines: list, use_synctime: bool = False, save_lines_to_txt=False, cut_out_outer_fixations=False603):604 """trial should be dict of line numbers of trials.605 lines should be list of lines from .asc file."""606 607 if use_synctime and "synctime" in trial:608 idx0, idxend = trial["synctime"] + 1, trial["trial_result_idx"]609 else:610 idx0, idxend = trial["start_idx"], trial["end_idx"]611 612 line_dicts = []613 fixations_dicts = []614 blink_started = False615 616 fixation_started = False617 efix_count = 0618 sfix_count = 0619 sblink_count = 0620 621 if save_lines_to_txt:622 with open("Lines_plus500.txt", "w") as f:623 f.writelines(lines[idx0 - 500 : idxend + 500])624 eye_to_use = "R"625 for l in lines[idx0 : idxend + 1]:626 if "EFIX R" in l:627 eye_to_use = "R"628 break629 elif "EFIX L" in l:630 eye_to_use = "L"631 break632 for l in lines[idx0 : idxend + 1]:633 parts = [x.strip() for x in l.split("\t")]634 if f"EFIX {eye_to_use}" in l:635 efix_count += 1636 if fixation_started:637 if parts[1] == "." and parts[2] == ".":638 continue639 fixations_dicts.append(640 {641 "start_time": float(parts[0].split()[-1].strip()),642 "end_time": float(parts[1].strip()),643 "duration": float(parts[2].strip()),644 "x": float(parts[3].strip()),645 "y": float(parts[4].strip()),646 "pupil_size": float(parts[5].strip()),647 }648 )649 if len(fixations_dicts) >= 2:650 assert (651 fixations_dicts[-1]["start_time"] > fixations_dicts[-2]["start_time"]652 ), "start times not in order"653 fixation_started = False654 655 elif f"SFIX {eye_to_use}" in l:656 sfix_count += 1657 fixation_started = True658 elif f"SBLINK {eye_to_use}" in l:659 sblink_count += 1660 blink_started = True661 if not blink_started and not any([True for x in event_strs if x in l]):662 if len(parts) < 3 or (parts[1] == "." and parts[2] == "."):663 continue664 line_dicts.append(665 {666 "idx": float(parts[0].strip()),667 "x": float(parts[1].strip()),668 "y": float(parts[2].strip()),669 "p": float(parts[3].strip()),670 }671 )672 673 elif f"EBLINK {eye_to_use}" in l:674 blink_started = False675 676 df = pd.DataFrame(line_dicts)677 dffix = pd.DataFrame(fixations_dicts)678 if len(fixations_dicts) > 0:679 dffix["corrected_start_time"] = dffix.start_time - trial["trial_start_time"]680 dffix["corrected_end_time"] = dffix.end_time - trial["trial_start_time"]681 dffix["fix_duration"] = dffix.corrected_end_time.values - dffix.corrected_start_time.values682 assert all(np.diff(dffix["corrected_start_time"]) > 0), "start times not in order"683 else:684 df, pd.DataFrame(), trial685 686 if cut_out_outer_fixations:687 dffix = dffix[(dffix.x > -10) & (dffix.y > -10) & (dffix.x < 1050) & (dffix.y < 800)]688 trial["efix_count"] = efix_count689 trial["eye_to_use"] = eye_to_use690 trial["sfix_count"] = sfix_count691 trial["sblink_count"] = sblink_count692 return df, dffix, trial693 694 695def get_save_path(fpath, fname_ending):696 save_path = gradio_plots.joinpath(f"{fpath.stem}_{fname_ending}.png")697 return save_path698 699 700def save_im_load_convert(fpath, fig, fname_ending, mode):701 save_path = get_save_path(fpath, fname_ending)702 fig.savefig(save_path)703 im = Image.open(save_path).convert(mode)704 im.save(save_path)705 return im706 707 708def get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, dffix=None, prefix="word"):709 fig = plt.figure(figsize=(screen_res[0] / dpi, screen_res[1] / dpi), dpi=dpi)710 ax = plt.Axes(fig, [0.0, 0.0, 1.0, 1.0])711 ax.set_axis_off()712 if dffix is not None:713 ax.set_ylim((dffix.y.min(), dffix.y.max()))714 ax.set_xlim((dffix.x.min(), dffix.x.max()))715 else:716 ax.set_ylim((words_df[f"{prefix}_y_center"].min() - y_margin, words_df[f"{prefix}_y_center"].max() + y_margin))717 ax.set_xlim((words_df[f"{prefix}_x_center"].min() - x_margin, words_df[f"{prefix}_x_center"].max() + x_margin))718 ax.invert_yaxis()719 fig.add_axes(ax)720 return fig, ax721 722 723def plot_text_boxes_fixations(724 fpath,725 dpi,726 screen_res,727 data_dir_sub,728 set_font_size: bool,729 font_size: int,730 use_words: bool,731 save_channel_repeats: bool,732 save_combo_grey_and_rgb: bool,733 dffix=None,734 trial=None,735):736 if isinstance(fpath, str):737 fpath = pl.Path(fpath)738 if use_words:739 prefix = "word"740 else:741 prefix = "char"742 if dffix is None:743 dffix = pd.read_csv(fpath)744 if trial is None:745 json_fpath = str(fpath).replace("_fixations.csv", "_trial.json")746 with open(json_fpath, "r") as f:747 trial = json.load(f)748 words_df = pd.DataFrame(trial[f"{prefix}s_list"])749 x_right = words_df[f"{prefix}_xmin"]750 x_left = words_df[f"{prefix}_xmax"]751 y_top = words_df[f"{prefix}_ymax"]752 y_bottom = words_df[f"{prefix}_ymin"]753 754 if f"{prefix}_x_center" not in words_df.columns:755 words_df[f"{prefix}_x_center"] = (words_df[f"{prefix}_xmax"] - words_df[f"{prefix}_xmin"]) / 2 + words_df[756 f"{prefix}_xmin"757 ]758 words_df[f"{prefix}_y_center"] = (words_df[f"{prefix}_ymax"] - words_df[f"{prefix}_ymin"]) / 2 + words_df[759 f"{prefix}_ymin"760 ]761 762 x_margin = words_df[f"{prefix}_x_center"].mean() / 8763 y_margin = words_df[f"{prefix}_y_center"].mean() / 4764 times = dffix.corrected_start_time - dffix.corrected_start_time.min()765 times = times / times.max()766 times = np.linspace(0.25, 1, len(times))767 768 if set_font_size:769 font = "monospace"770 else:771 font_size = trial["font_size"] * 27 // dpi772 773 font_props = FontProperties(family=font, style="normal", size=font_size)774 if save_combo_grey_and_rgb:775 fig, ax = get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, prefix=prefix)776 ax.scatter(dffix.x, dffix.y, alpha=times, facecolor="b")777 for idx in range(len(x_left)):778 xdiff = x_right[idx] - x_left[idx]779 ydiff = y_top[idx] - y_bottom[idx]780 rect = patches.Rectangle(781 (x_left[idx] - 1, y_bottom[idx] - 1),782 xdiff,783 ydiff,784 alpha=0.9,785 linewidth=0.8,786 edgecolor="r",787 facecolor="none",788 ) # seems to need one pixel offset789 ax.text(790 words_df[f"{prefix}_x_center"][idx],791 words_df[f"{prefix}_y_center"][idx],792 words_df[prefix][idx],793 horizontalalignment="center",794 verticalalignment="center",795 fontproperties=font_props,796 color="g",797 )798 ax.add_patch(rect)799 fname_ending = f"{prefix}s_combo_rgb"800 words_combo_rgb_im = save_im_load_convert(fpath, fig, fname_ending, "RGB")801 plt.close("all")802 803 fig, ax = get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, prefix=prefix)804 805 ax.scatter(dffix.x, dffix.y, facecolor="k", alpha=times)806 for idx in range(len(x_left)):807 xdiff = x_right[idx] - x_left[idx]808 ydiff = y_top[idx] - y_bottom[idx]809 rect = patches.Rectangle(810 (x_left[idx] - 1, y_bottom[idx] - 1),811 xdiff,812 ydiff,813 alpha=0.9,814 linewidth=0.8,815 edgecolor="k",816 facecolor="none",817 ) # seems to need one pixel offset818 ax.text(819 words_df[f"{prefix}_x_center"][idx],820 words_df[f"{prefix}_y_center"][idx],821 words_df[prefix][idx],822 horizontalalignment="center",823 verticalalignment="center",824 fontproperties=font_props,825 )826 ax.add_patch(rect)827 fname_ending = f"{prefix}s_combo_grey"828 words_combo_grey_im = save_im_load_convert(fpath, fig, fname_ending, "L")829 plt.close("all")830 831 fig, ax = get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, prefix=prefix)832 833 ax.scatter(words_df[f"{prefix}_x_center"], words_df[f"{prefix}_y_center"], s=1, facecolor="k", alpha=0.01)834 for idx in range(len(x_left)):835 ax.text(836 words_df[f"{prefix}_x_center"][idx],837 words_df[f"{prefix}_y_center"][idx],838 words_df[prefix][idx],839 horizontalalignment="center",840 verticalalignment="center",841 fontproperties=font_props,842 )843 fname_ending = f"{prefix}s_grey"844 words_grey_im = save_im_load_convert(fpath, fig, fname_ending, "L")845 846 plt.close("all")847 fig, ax = get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, prefix=prefix)848 849 ax.scatter(words_df[f"{prefix}_x_center"], words_df[f"{prefix}_y_center"], s=1, facecolor="k", alpha=0.1)850 for idx in range(len(x_left)):851 xdiff = x_right[idx] - x_left[idx]852 ydiff = y_top[idx] - y_bottom[idx]853 rect = patches.Rectangle(854 (x_left[idx] - 1, y_bottom[idx] - 1), xdiff, ydiff, alpha=0.9, linewidth=1, edgecolor="k", facecolor="grey"855 ) # seems to need one pixel offset856 ax.add_patch(rect)857 fname_ending = f"{prefix}_boxes_grey"858 word_boxes_grey_im = save_im_load_convert(fpath, fig, fname_ending, "L")859 860 plt.close("all")861 862 fig, ax = get_fig_ax(screen_res, dpi, words_df, x_margin, y_margin, prefix=prefix)863 864 ax.scatter(dffix.x, dffix.y, facecolor="k", alpha=times)865 fname_ending = "fix_scatter_grey"866 fix_scatter_grey_im = save_im_load_convert(fpath, fig, fname_ending, "L")867 868 plt.close("all")869 870 arr_combo = np.stack(871 [872 np.asarray(words_grey_im),873 np.asarray(word_boxes_grey_im),874 np.asarray(fix_scatter_grey_im),875 ],876 axis=2,877 )878 879 im_combo = Image.fromarray(arr_combo)880 fname_ending = f"{prefix}s_channel_sep"881 882 save_path = get_save_path(fpath, fname_ending)883 print(f"save_path for im combo is {save_path}")884 im_combo.save(fpath)885 886 if save_channel_repeats:887 arr_combo = np.stack([np.asarray(words_grey_im)] * 3, axis=2)888 im_combo = Image.fromarray(arr_combo)889 fname_ending = f"{prefix}s_channel_repeat"890 891 save_path = get_save_path(fpath, fname_ending)892 im_combo.save(save_path)893 894 arr_combo = np.stack([np.asarray(word_boxes_grey_im)] * 3, axis=2)895 896 im_combo = Image.fromarray(arr_combo)897 fname_ending = f"{prefix}boxes_channel_repeat"898 899 save_path = get_save_path(fpath, fname_ending)900 im_combo.save(save_path)901 902 arr_combo = np.stack([np.asarray(fix_scatter_grey_im)] * 3, axis=2)903 904 im_combo = Image.fromarray(arr_combo)905 fname_ending = "fix_channel_repeat"906 907 save_path = get_save_path(fpath, fname_ending)908 im_combo.save(save_path)909 910 911def add_line_overlaps_to_sample(trial, sample):912 char_df = pd.DataFrame(trial["chars_list"])913 line_overlaps = []914 for arr in sample:915 y_val = arr[1]916 line_overlap = t.tensor(-1, dtype=t.float32)917 for idx, (x1, x2) in enumerate(zip(char_df.char_ymin.unique(), char_df.char_ymax.unique())):918 if x1 <= y_val <= x2:919 line_overlap = t.tensor(idx, dtype=t.float32)920 break921 line_overlaps.append(line_overlap)922 line_olaps_tensor = t.stack(line_overlaps, dim=0)923 sample = t.cat([sample, line_olaps_tensor.unsqueeze(1)], dim=1)924 return sample925 926 927def norm_coords_by_letter_min_x_y(928 sample_idx: int,929 trialslist: list,930 samplelist: list,931 chars_center_coords_list: list = None,932):933 chars_df = pd.DataFrame(trialslist[sample_idx]["chars_list"])934 trialslist[sample_idx]["x_char_unique"] = chars_df.char_xmin.unique()935 936 min_x_chars = chars_df.char_xmin.min()937 min_y_chars = chars_df.char_ymin.min()938 939 norm_vector_substract = t.zeros(940 (1, samplelist[sample_idx].shape[1]), dtype=samplelist[sample_idx].dtype, device=samplelist[sample_idx].device941 )942 norm_vector_substract[0, 0] = norm_vector_substract[0, 0] + 1 * min_x_chars943 norm_vector_substract[0, 1] = norm_vector_substract[0, 1] + 1 * min_y_chars944 945 samplelist[sample_idx] = samplelist[sample_idx] - norm_vector_substract946 947 if chars_center_coords_list is not None:948 norm_vector_substract = norm_vector_substract.squeeze(0)[:2]949 if chars_center_coords_list[sample_idx].shape[-1] == norm_vector_substract.shape[-1] * 2:950 chars_center_coords_list[sample_idx][:, :2] -= norm_vector_substract951 chars_center_coords_list[sample_idx][:, 2:] -= norm_vector_substract952 else:953 chars_center_coords_list[sample_idx] -= norm_vector_substract954 return trialslist, samplelist, chars_center_coords_list955 956 957def norm_coords_by_letter_positions(958 sample_idx: int,959 trialslist: list,960 samplelist: list,961 meanlist: list = None,962 stdlist: list = None,963 return_mean_std_lists=False,964 norm_by_char_averages=False,965 chars_center_coords_list: list = None,966 add_normalised_values_as_features=False,967):968 chars_df = pd.DataFrame(trialslist[sample_idx]["chars_list"])969 trialslist[sample_idx]["x_char_unique"] = chars_df.char_xmin.unique()970 971 min_x_chars = chars_df.char_xmin.min()972 max_x_chars = chars_df.char_xmax.max()973 974 norm_vector_multi = t.ones(975 (1, samplelist[sample_idx].shape[1]), dtype=samplelist[sample_idx].dtype, device=samplelist[sample_idx].device976 )977 if norm_by_char_averages:978 chars_list = trialslist[sample_idx]["chars_list"]979 char_widths = np.asarray([x["char_xmax"] - x["char_xmin"] for x in chars_list])980 char_heights = np.asarray([x["char_ymax"] - x["char_ymin"] for x in chars_list])981 char_widths_average = np.mean(char_widths[char_widths > 0])982 char_heights_average = np.mean(char_heights[char_heights > 0])983 984 norm_vector_multi[0, 0] = norm_vector_multi[0, 0] * char_widths_average985 norm_vector_multi[0, 1] = norm_vector_multi[0, 1] * char_heights_average986 987 else:988 line_height = min(np.unique(trialslist[sample_idx]["line_heights"]))989 line_width = max_x_chars - min_x_chars990 norm_vector_multi[0, 0] = norm_vector_multi[0, 0] * line_width991 norm_vector_multi[0, 1] = norm_vector_multi[0, 1] * line_height992 assert ~t.any(t.isnan(norm_vector_multi)), "Nan found in char norming vector"993 994 norm_vector_multi = norm_vector_multi.squeeze(0)995 if add_normalised_values_as_features:996 norm_vector_multi = norm_vector_multi[norm_vector_multi != 1]997 normed_features = samplelist[sample_idx][:, : norm_vector_multi.shape[0]] / norm_vector_multi998 samplelist[sample_idx] = t.cat([samplelist[sample_idx], normed_features], dim=1)999 else:1000 samplelist[sample_idx] = samplelist[sample_idx] / norm_vector_multi # in case time or pupil size is included1001 if chars_center_coords_list is not None:1002 norm_vector_multi = norm_vector_multi[:2]1003 if chars_center_coords_list[sample_idx].shape[-1] == norm_vector_multi.shape[-1] * 2:1004 chars_center_coords_list[sample_idx][:, :2] /= norm_vector_multi1005 chars_center_coords_list[sample_idx][:, 2:] /= norm_vector_multi1006 else:1007 chars_center_coords_list[sample_idx] /= norm_vector_multi1008 if return_mean_std_lists:1009 mean_val = samplelist[sample_idx].mean(axis=0).cpu().numpy()1010 meanlist.append(mean_val)1011 std_val = samplelist[sample_idx].std(axis=0).cpu().numpy()1012 stdlist.append(std_val)1013 assert ~any(np.isnan(mean_val)), "Nan found in mean_val"1014 assert ~any(np.isnan(mean_val)), "Nan found in std_val"1015 1016 return trialslist, samplelist, meanlist, stdlist, chars_center_coords_list1017 return trialslist, samplelist, chars_center_coords_list1018 1019 1020def remove_compile_from_model(model):1021 if hasattr(model.project, "_orig_mod"):1022 model.project = model.project._orig_mod1023 model.chars_conv = model.chars_conv._orig_mod1024 model.chars_classifier = model.chars_classifier._orig_mod1025 model.layer_norm_in = model.layer_norm_in._orig_mod1026 model.bert_model = model.bert_model._orig_mod1027 model.linear = model.linear._orig_mod1028 else:1029 print(f"remove_compile_from_model not done since model.project {model.project} has no orig_mod")1030 return model1031 1032 1033def remove_compile_from_dict(state_dict):1034 for key in list(state_dict.keys()):1035 newkey = key.replace("._orig_mod.", ".")1036 state_dict[newkey] = state_dict.pop(key)1037 return state_dict1038 1039 1040def add_text_to_ax(1041 chars_list,1042 ax,1043 font_to_use="DejaVu Sans Mono",1044 fontsize=21,1045 prefix="char",1046 plot_boxes=True,1047 plot_text=True,1048 box_annotations=None,1049):1050 font_props = FontProperties(family=font_to_use, style="normal", size=fontsize)1051 if not plot_boxes and not plot_text:1052 return None1053 if box_annotations is None:1054 enum = chars_list1055 else:1056 enum = zip(chars_list, box_annotations)1057 for v in enum:1058 if box_annotations is not None:1059 v, annot_text = v1060 x0, y0 = v[f"{prefix}_xmin"], v[f"{prefix}_ymin"]1061 xdiff, ydiff = v[f"{prefix}_xmax"] - v[f"{prefix}_xmin"], v[f"{prefix}_ymax"] - v[f"{prefix}_ymin"]1062 if plot_text:1063 ax.text(1064 v[f"{prefix}_x_center"],1065 v[f"{prefix}_y_center"],1066 v[prefix],1067 horizontalalignment="center",1068 verticalalignment="center",1069 fontproperties=font_props,1070 )1071 if plot_boxes:1072 ax.add_patch(Rectangle((x0, y0), xdiff, ydiff, edgecolor="grey", facecolor="none", lw=0.8, alpha=0.4))1073 if box_annotations is not None:1074 ax.annotate(1075 str(annot_text),1076 (x0 + xdiff / 2, y0),1077 horizontalalignment="center",1078 verticalalignment="center",1079 fontproperties=FontProperties(family=font_to_use, style="normal", size=fontsize / 1.5),1080 )1081 1082 1083def plot_fixations_and_text(1084 dffix: pd.DataFrame,1085 trial: dict,1086 plot_prefix="chars_",1087 show=False,1088 returnfig=False,1089 save=False,1090 savelocation="plot.png",1091 font_to_use="DejaVu Sans Mono",1092 fontsize=20,1093 plot_classic=True,1094 plot_boxes=True,1095 plot_text=True,1096 fig_size=(14, 8),1097 dpi=300,1098 turn_axis_on=True,1099 algo_choice="slice",1100):1101 fig, ax = plt.subplots(1, 1, figsize=fig_size, tight_layout=True, dpi=dpi)1102 if f"{plot_prefix}list" in trial.keys():1103 add_text_to_ax(1104 trial[f"{plot_prefix}list"],1105 ax,1106 font_to_use,1107 fontsize=fontsize,1108 prefix=plot_prefix[:-2],1109 plot_boxes=plot_boxes,1110 plot_text=plot_text,1111 )1112 ax.plot(dffix.x, dffix.y, "kX", label="Raw Fixations", alpha=0.9)1113 1114 if plot_classic and f"line_num_{algo_choice}" in dffix.columns:1115 ax.scatter(1116 dffix.x,1117 dffix[f"y_{algo_choice}"],1118 marker="*",1119 color="tab:green",1120 label=f"{algo_choice} Prediction",1121 alpha=0.9,1122 )1123 for x_before, y_before, x_after, y_after in zip(1124 dffix.x.values, dffix[f"y_{algo_choice}"].values, dffix.x, dffix.y1125 ):1126 arr_delta_x = x_after - x_before1127 arr_delta_y = y_after - y_before1128 ax.arrow(x_before, y_before, arr_delta_x, arr_delta_y, color="tab:green", alpha=0.6)1129 ax.set_ylabel("y (pixel)")1130 ax.set_xlabel("x (pixel)")1131 1132 ax.invert_yaxis()1133 ax.legend(bbox_to_anchor=(1, 1), loc="upper left")1134 if not turn_axis_on:1135 ax.axis("off")1136 if save:1137 plt.savefig(savelocation, dpi=dpi)1138 if show:1139 plt.show()1140 if returnfig:1141 return fig1142 else:1143 plt.close()1144 return None1145 1146 1147def make_folders(gradio_temp_folder, gradio_temp_unzipped_folder, gradio_plots):1148 gradio_temp_folder.mkdir(exist_ok=True)1149 gradio_temp_unzipped_folder.mkdir(exist_ok=True)1150 gradio_plots.mkdir(exist_ok=True)1151 return 01152 1153 1154def get_classic_cfg(fname):1155 with open(fname, "r") as f:1156 jsonsstring = f.read()1157 classic_algos_cfg = json.loads(jsonsstring)1158 classic_algos_cfg["slice"] = classic_algos_cfg["slice"]1159 classic_algos_cfg = classic_algos_cfg1160 return classic_algos_cfg1161 1162 1163def find_and_load_model(model_date="20240104-223349"):1164 model_cfg_file = list(DIST_MODELS_FOLDER.glob(f"*{model_date}*.yaml"))1165 if len(model_cfg_file) == 0:1166 if "logger" in st.session_state:1167 st.session_state["logger"].warning(f"No model cfg yaml found for {model_date}")1168 return None, None1169 model_cfg_file = model_cfg_file[0]1170 with open(model_cfg_file) as f:1171 model_cfg = yaml.safe_load(f)1172 1173 model_cfg["system_type"] = "linux"1174 model_file = list(pl.Path("models").glob(f"*{model_date}*.ckpt"))[0]1175 model = load_model(model_file, model_cfg)1176 1177 return model, model_cfg1178 1179 1180def load_model(model_file, cfg):1181 try:1182 model_loaded = t.load(model_file, map_location="cpu")1183 if "hyper_parameters" in model_loaded.keys():1184 model_cfg_temp = model_loaded["hyper_parameters"]["cfg"]1185 else:1186 model_cfg_temp = cfg1187 model_state_dict = model_loaded["state_dict"]1188 except Exception as e:1189 if "logger" in st.session_state:1190 st.session_state["logger"].warning(e)1191 if "logger" in st.session_state:1192 st.session_state["logger"].warning(f"Failed to load {model_file}")1193 return None1194 model = LitModel(1195 [1, 500, 3],1196 model_cfg_temp["hidden_dim_bert"],1197 model_cfg_temp["num_attention_heads"],1198 model_cfg_temp["n_layers_BERT"],1199 model_cfg_temp["loss_function"],1200 1e-4,