exbert-project/exbert
185
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 