CoolFace
Apppublic

bugroup/Eye_Tracking_Drift_Correction

sourceHugging Faceupdated 3y agoView on Hugging Face
3likes
utils.py2016 linesDownload Raw Back to root
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,

Showing the first 1,200 of 2016 lines. Download the file for the rest.