Team Ai
Apppublic

exbert-project/exbert

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
185likes
uiConfig.ts283 linesDownload Raw Back to ts
1import * as tp from "./etc/types"2import * as x_ from "./etc/_Tools"3import * as _ from "lodash"4import * as R from 'ramda'5import { URLHandler } from "./etc/URLHandler";6 7const falsey = val => (new Set(['false', 0, "no", false, null, ""])).has(val)8const truthy = val => !falsey(val)9const toNumber = x => +x;10 11 12type InspectorOptions = "context" | "embeddings" | null13 14// Must be optional params for initializations15interface URLParameters {16    sentence?: string17    model?: string18    modelKind?: string19    layer?: number20    heads?: number[]21    threshold?: number22    tokenInd?: number | 'null'23    tokenSide?: tp.SideOptions24    maskInds?: number[]25    hideClsSep?: boolean26}27 28export class UIConfig {29 30    private _conf: URLParameters = {}31    private _headSet: Set<number>;32    attType: "aa"33    _nHeads: number | null;34    _nLayers: number | null;35    private _token: tp.TokenEvent;36 37    constructor() {38        this._nHeads = 12; 39        this._nLayers = null;40        this.attType = 'aa'41        this.fromURL()42        this.toURL(false)43    }44 45    toURL(updateHistory = false) {46        URLHandler.updateUrl(this._conf, updateHistory)47    }48 49 50    fromURL() {51        const params = URLHandler.parameters52 53        this._conf = {54            model: params['model'] || 'bert-base-cased',55            modelKind: params['modelKind'] || tp.ModelKind.Bidirectional,56            sentence: params['sentence'] || "The girl ran to a local pub to escape the din of her city.",57            layer: params['layer'] || 1,58            heads: this._initHeads(params['heads']),59            threshold: params['threshold'] || 0.7,60            tokenInd: params['tokenInd'] || null,61            tokenSide: params['tokenSide'] || null,62            maskInds: params['maskInds'] || [9],63            hideClsSep: truthy(params['hideClsSep']) || true,64        }65 66        this._token = { side: this._conf.tokenSide, ind: this._conf.tokenInd }67 68    }69 70    private _initHeads(v: number[] | null) {71        if (v == null || v.length < 1) {72            this.selectAllHeads()73        }74        else {75            console.log(this.headSet(new Set(v))._conf.heads);76        }77 78        return this.heads()79    }80 81    nHeads(): number82    nHeads(val: number): this83    nHeads(val?) {84        if (val == null) return this._nHeads85        this._nHeads = val86        return this87    }88 89    nLayers(): number90    nLayers(val: number): this91    nLayers(val?) {92        if (val == null) return this._nLayers93        this._nLayers = val94        return this95    }96 97    toggleSelectAllHeads() {98        if (this.heads().length == 0) {99            this.selectAllHeads()100        }101        else {102            this.selectNoHeads()103        }104    }105 106    selectAllHeads() {107        this.headSet(new Set(_.range(0, this._nHeads)))108    }109 110    selectNoHeads() {111        this.headSet(new Set([]))112    }113 114    toggleHead(head: number): tp.Toggled {115        let out;116        if (this.headSet().has(head)) {117            this.headSet().delete(head);118            out = tp.Toggled.REMOVED119        }120        else {121            this.headSet().add(head);122            out = tp.Toggled.ADDED123        }124 125        // Set through setter function to ensure url is updated126        this.headSet(this.headSet()); // I hate mutable datastructures... This is confusing.127 128        return out129    }130 131    toggleToken(e: tp.TokenEvent): this {132        const picker = R.pick(['ind', 'side'])133        const compareEvent = picker(e)134        const compareToken = picker(this.token())135 136        if (R.equals(compareToken, compareEvent)) {137            this.rmToken();138        }139        else {140            this.token(e);141        }142        return this;143    }144 145    token(): tp.TokenEvent;146    token(val: tp.TokenEvent): this;147    token(val?: tp.TokenEvent) {148        if (val == null)149            return this._token150 151        this._token = val;152        this._conf.tokenInd = val.ind;153        this._conf.tokenSide = val.side;154        this.toURL();155 156        return this157    }158 159    hasToken() {160        const conf = this._conf161        const actuallyNull = ((conf.tokenInd == null) && (conf.tokenSide == null))162        const strNull = (conf.tokenInd == "null")163        return (!actuallyNull) && (!strNull)164    }165 166    rmToken() {167        this.token({ ind: null, side: null });168        return this169    }170 171    sentence(): string;172    sentence(val: string): this;173    sentence(val?) {174        if (val == null)175            return this._conf.sentence176 177        this._conf.sentence = val178        this.toURL(true)179        return this180    }181 182    threshold(): number;183    threshold(val: number): this;184    threshold(val?) {185        if (val == null) return this._conf.threshold;186 187        this._conf.threshold = val;188        this.toURL();189        return this;190    }191 192    heads(): number[] {193        return this._conf.heads194    }195 196    layer(): number197    layer(val: number): this198    layer(val?) {199        if (val == null)200            return this._conf.layer201 202        this._conf.layer = val;203        this.toURL();204        return this205    }206 207    headSet(): Set<number>;208    headSet(val: Set<number>): this;209    headSet(val?) {210        if (val == null) {211            return this._headSet;212        }213 214        this._headSet = val;215        this._conf.heads = x_.set2SortedArray(this._headSet)216        this.toURL();217        return this218    }219 220    maskInds(): number[];221    maskInds(val: number[]): this;222    maskInds(val?) {223        if (val == null) return this._conf.maskInds;224 225        this._conf.maskInds = val;226        this.toURL();227        return this;228    }229 230    hideClsSep(): boolean;231    hideClsSep(val: boolean): this;232    hideClsSep(val?) {233        if (val == null) return this._conf.hideClsSep;234 235        this._conf.hideClsSep = truthy(val);236        this.toURL();237        return this;238    }239 240    model(): string;241    model(val: string): this;242    model(val?) {243        if (val == null) return this._conf.model244        this._conf.model = val245        this.toURL();246        return this247    }248 249    modelKind(): string;250    modelKind(val: string): this;251    modelKind(val?) {252        if (val == null) return this._conf.modelKind253        this._conf.modelKind = val254        this.toURL();255        return this256    }257 258    /**259     * Return the offset needed for the modelKind in the configuration260     */261    get offset() {262        switch (this.modelKind()) {263            case tp.ModelKind.Bidirectional: {264                return 0265            }266            case tp.ModelKind.Autoregressive: {267                return 0268            }269            default: {270                return 0271            }272        }273    }274 275    get showNext() {276        return this.modelKind() == tp.ModelKind.Autoregressive ? true : false277    }278 279    get matchHistogramDescription() {280        return this.modelKind() == tp.ModelKind.Autoregressive ? "Next" : "Matched"281    }282}283