ai-zero-shot-classifier
Version:
🧠powerful JavaScript library that leverages advanced AI embeddings to perform zero-shot text classification. Whether you're dealing with unlabelled data or seeking to classify text against dynamic and user-defined labels, this library provides a seamles
369 lines (358 loc) • 27 kB
JavaScript
"use strict";
Object.defineProperty(exports, "__esModule", {
value: true
});
exports.default = void 0;
var _lodash = _interopRequireDefault(require("lodash.chunk"));
var _pMap = _interopRequireDefault(require("../utils/p-map"));
var _providers = require("../providers");
var _similarityFunctions = require("../utils/similarityFunctions");
function _interopRequireDefault(e) { return e && e.__esModule ? e : { default: e }; }
function _typeof(o) { "@babel/helpers - typeof"; return _typeof = "function" == typeof Symbol && "symbol" == typeof Symbol.iterator ? function (o) { return typeof o; } : function (o) { return o && "function" == typeof Symbol && o.constructor === Symbol && o !== Symbol.prototype ? "symbol" : typeof o; }, _typeof(o); }
function _toConsumableArray(r) { return _arrayWithoutHoles(r) || _iterableToArray(r) || _unsupportedIterableToArray(r) || _nonIterableSpread(); }
function _nonIterableSpread() { throw new TypeError("Invalid attempt to spread non-iterable instance.\nIn order to be iterable, non-array objects must have a [Symbol.iterator]() method."); }
function _iterableToArray(r) { if ("undefined" != typeof Symbol && null != r[Symbol.iterator] || null != r["@@iterator"]) return Array.from(r); }
function _arrayWithoutHoles(r) { if (Array.isArray(r)) return _arrayLikeToArray(r); }
function _regeneratorRuntime() { "use strict"; /*! regenerator-runtime -- Copyright (c) 2014-present, Facebook, Inc. -- license (MIT): https://github.com/facebook/regenerator/blob/main/LICENSE */ _regeneratorRuntime = function _regeneratorRuntime() { return e; }; var t, e = {}, r = Object.prototype, n = r.hasOwnProperty, o = Object.defineProperty || function (t, e, r) { t[e] = r.value; }, i = "function" == typeof Symbol ? Symbol : {}, a = i.iterator || "@@iterator", c = i.asyncIterator || "@@asyncIterator", u = i.toStringTag || "@@toStringTag"; function define(t, e, r) { return Object.defineProperty(t, e, { value: r, enumerable: !0, configurable: !0, writable: !0 }), t[e]; } try { define({}, ""); } catch (t) { define = function define(t, e, r) { return t[e] = r; }; } function wrap(t, e, r, n) { var i = e && e.prototype instanceof Generator ? e : Generator, a = Object.create(i.prototype), c = new Context(n || []); return o(a, "_invoke", { value: makeInvokeMethod(t, r, c) }), a; } function tryCatch(t, e, r) { try { return { type: "normal", arg: t.call(e, r) }; } catch (t) { return { type: "throw", arg: t }; } } e.wrap = wrap; var h = "suspendedStart", l = "suspendedYield", f = "executing", s = "completed", y = {}; function Generator() {} function GeneratorFunction() {} function GeneratorFunctionPrototype() {} var p = {}; define(p, a, function () { return this; }); var d = Object.getPrototypeOf, v = d && d(d(values([]))); v && v !== r && n.call(v, a) && (p = v); var g = GeneratorFunctionPrototype.prototype = Generator.prototype = Object.create(p); function defineIteratorMethods(t) { ["next", "throw", "return"].forEach(function (e) { define(t, e, function (t) { return this._invoke(e, t); }); }); } function AsyncIterator(t, e) { function invoke(r, o, i, a) { var c = tryCatch(t[r], t, o); if ("throw" !== c.type) { var u = c.arg, h = u.value; return h && "object" == _typeof(h) && n.call(h, "__await") ? e.resolve(h.__await).then(function (t) { invoke("next", t, i, a); }, function (t) { invoke("throw", t, i, a); }) : e.resolve(h).then(function (t) { u.value = t, i(u); }, function (t) { return invoke("throw", t, i, a); }); } a(c.arg); } var r; o(this, "_invoke", { value: function value(t, n) { function callInvokeWithMethodAndArg() { return new e(function (e, r) { invoke(t, n, e, r); }); } return r = r ? r.then(callInvokeWithMethodAndArg, callInvokeWithMethodAndArg) : callInvokeWithMethodAndArg(); } }); } function makeInvokeMethod(e, r, n) { var o = h; return function (i, a) { if (o === f) throw Error("Generator is already running"); if (o === s) { if ("throw" === i) throw a; return { value: t, done: !0 }; } for (n.method = i, n.arg = a;;) { var c = n.delegate; if (c) { var u = maybeInvokeDelegate(c, n); if (u) { if (u === y) continue; return u; } } if ("next" === n.method) n.sent = n._sent = n.arg;else if ("throw" === n.method) { if (o === h) throw o = s, n.arg; n.dispatchException(n.arg); } else "return" === n.method && n.abrupt("return", n.arg); o = f; var p = tryCatch(e, r, n); if ("normal" === p.type) { if (o = n.done ? s : l, p.arg === y) continue; return { value: p.arg, done: n.done }; } "throw" === p.type && (o = s, n.method = "throw", n.arg = p.arg); } }; } function maybeInvokeDelegate(e, r) { var n = r.method, o = e.iterator[n]; if (o === t) return r.delegate = null, "throw" === n && e.iterator.return && (r.method = "return", r.arg = t, maybeInvokeDelegate(e, r), "throw" === r.method) || "return" !== n && (r.method = "throw", r.arg = new TypeError("The iterator does not provide a '" + n + "' method")), y; var i = tryCatch(o, e.iterator, r.arg); if ("throw" === i.type) return r.method = "throw", r.arg = i.arg, r.delegate = null, y; var a = i.arg; return a ? a.done ? (r[e.resultName] = a.value, r.next = e.nextLoc, "return" !== r.method && (r.method = "next", r.arg = t), r.delegate = null, y) : a : (r.method = "throw", r.arg = new TypeError("iterator result is not an object"), r.delegate = null, y); } function pushTryEntry(t) { var e = { tryLoc: t[0] }; 1 in t && (e.catchLoc = t[1]), 2 in t && (e.finallyLoc = t[2], e.afterLoc = t[3]), this.tryEntries.push(e); } function resetTryEntry(t) { var e = t.completion || {}; e.type = "normal", delete e.arg, t.completion = e; } function Context(t) { this.tryEntries = [{ tryLoc: "root" }], t.forEach(pushTryEntry, this), this.reset(!0); } function values(e) { if (e || "" === e) { var r = e[a]; if (r) return r.call(e); if ("function" == typeof e.next) return e; if (!isNaN(e.length)) { var o = -1, i = function next() { for (; ++o < e.length;) if (n.call(e, o)) return next.value = e[o], next.done = !1, next; return next.value = t, next.done = !0, next; }; return i.next = i; } } throw new TypeError(_typeof(e) + " is not iterable"); } return GeneratorFunction.prototype = GeneratorFunctionPrototype, o(g, "constructor", { value: GeneratorFunctionPrototype, configurable: !0 }), o(GeneratorFunctionPrototype, "constructor", { value: GeneratorFunction, configurable: !0 }), GeneratorFunction.displayName = define(GeneratorFunctionPrototype, u, "GeneratorFunction"), e.isGeneratorFunction = function (t) { var e = "function" == typeof t && t.constructor; return !!e && (e === GeneratorFunction || "GeneratorFunction" === (e.displayName || e.name)); }, e.mark = function (t) { return Object.setPrototypeOf ? Object.setPrototypeOf(t, GeneratorFunctionPrototype) : (t.__proto__ = GeneratorFunctionPrototype, define(t, u, "GeneratorFunction")), t.prototype = Object.create(g), t; }, e.awrap = function (t) { return { __await: t }; }, defineIteratorMethods(AsyncIterator.prototype), define(AsyncIterator.prototype, c, function () { return this; }), e.AsyncIterator = AsyncIterator, e.async = function (t, r, n, o, i) { void 0 === i && (i = Promise); var a = new AsyncIterator(wrap(t, r, n, o), i); return e.isGeneratorFunction(r) ? a : a.next().then(function (t) { return t.done ? t.value : a.next(); }); }, defineIteratorMethods(g), define(g, u, "Generator"), define(g, a, function () { return this; }), define(g, "toString", function () { return "[object Generator]"; }), e.keys = function (t) { var e = Object(t), r = []; for (var n in e) r.push(n); return r.reverse(), function next() { for (; r.length;) { var t = r.pop(); if (t in e) return next.value = t, next.done = !1, next; } return next.done = !0, next; }; }, e.values = values, Context.prototype = { constructor: Context, reset: function reset(e) { if (this.prev = 0, this.next = 0, this.sent = this._sent = t, this.done = !1, this.delegate = null, this.method = "next", this.arg = t, this.tryEntries.forEach(resetTryEntry), !e) for (var r in this) "t" === r.charAt(0) && n.call(this, r) && !isNaN(+r.slice(1)) && (this[r] = t); }, stop: function stop() { this.done = !0; var t = this.tryEntries[0].completion; if ("throw" === t.type) throw t.arg; return this.rval; }, dispatchException: function dispatchException(e) { if (this.done) throw e; var r = this; function handle(n, o) { return a.type = "throw", a.arg = e, r.next = n, o && (r.method = "next", r.arg = t), !!o; } for (var o = this.tryEntries.length - 1; o >= 0; --o) { var i = this.tryEntries[o], a = i.completion; if ("root" === i.tryLoc) return handle("end"); if (i.tryLoc <= this.prev) { var c = n.call(i, "catchLoc"), u = n.call(i, "finallyLoc"); if (c && u) { if (this.prev < i.catchLoc) return handle(i.catchLoc, !0); if (this.prev < i.finallyLoc) return handle(i.finallyLoc); } else if (c) { if (this.prev < i.catchLoc) return handle(i.catchLoc, !0); } else { if (!u) throw Error("try statement without catch or finally"); if (this.prev < i.finallyLoc) return handle(i.finallyLoc); } } } }, abrupt: function abrupt(t, e) { for (var r = this.tryEntries.length - 1; r >= 0; --r) { var o = this.tryEntries[r]; if (o.tryLoc <= this.prev && n.call(o, "finallyLoc") && this.prev < o.finallyLoc) { var i = o; break; } } i && ("break" === t || "continue" === t) && i.tryLoc <= e && e <= i.finallyLoc && (i = null); var a = i ? i.completion : {}; return a.type = t, a.arg = e, i ? (this.method = "next", this.next = i.finallyLoc, y) : this.complete(a); }, complete: function complete(t, e) { if ("throw" === t.type) throw t.arg; return "break" === t.type || "continue" === t.type ? this.next = t.arg : "return" === t.type ? (this.rval = this.arg = t.arg, this.method = "return", this.next = "end") : "normal" === t.type && e && (this.next = e), y; }, finish: function finish(t) { for (var e = this.tryEntries.length - 1; e >= 0; --e) { var r = this.tryEntries[e]; if (r.finallyLoc === t) return this.complete(r.completion, r.afterLoc), resetTryEntry(r), y; } }, catch: function _catch(t) { for (var e = this.tryEntries.length - 1; e >= 0; --e) { var r = this.tryEntries[e]; if (r.tryLoc === t) { var n = r.completion; if ("throw" === n.type) { var o = n.arg; resetTryEntry(r); } return o; } } throw Error("illegal catch attempt"); }, delegateYield: function delegateYield(e, r, n) { return this.delegate = { iterator: values(e), resultName: r, nextLoc: n }, "next" === this.method && (this.arg = t), y; } }, e; }
function asyncGeneratorStep(n, t, e, r, o, a, c) { try { var i = n[a](c), u = i.value; } catch (n) { return void e(n); } i.done ? t(u) : Promise.resolve(u).then(r, o); }
function _asyncToGenerator(n) { return function () { var t = this, e = arguments; return new Promise(function (r, o) { var a = n.apply(t, e); function _next(n) { asyncGeneratorStep(a, r, o, _next, _throw, "next", n); } function _throw(n) { asyncGeneratorStep(a, r, o, _next, _throw, "throw", n); } _next(void 0); }); }; }
function _slicedToArray(r, e) { return _arrayWithHoles(r) || _iterableToArrayLimit(r, e) || _unsupportedIterableToArray(r, e) || _nonIterableRest(); }
function _nonIterableRest() { throw new TypeError("Invalid attempt to destructure non-iterable instance.\nIn order to be iterable, non-array objects must have a [Symbol.iterator]() method."); }
function _unsupportedIterableToArray(r, a) { if (r) { if ("string" == typeof r) return _arrayLikeToArray(r, a); var t = {}.toString.call(r).slice(8, -1); return "Object" === t && r.constructor && (t = r.constructor.name), "Map" === t || "Set" === t ? Array.from(r) : "Arguments" === t || /^(?:Ui|I)nt(?:8|16|32)(?:Clamped)?Array$/.test(t) ? _arrayLikeToArray(r, a) : void 0; } }
function _arrayLikeToArray(r, a) { (null == a || a > r.length) && (a = r.length); for (var e = 0, n = Array(a); e < a; e++) n[e] = r[e]; return n; }
function _iterableToArrayLimit(r, l) { var t = null == r ? null : "undefined" != typeof Symbol && r[Symbol.iterator] || r["@@iterator"]; if (null != t) { var e, n, i, u, a = [], f = !0, o = !1; try { if (i = (t = t.call(r)).next, 0 === l) { if (Object(t) !== t) return; f = !1; } else for (; !(f = (e = i.call(t)).done) && (a.push(e.value), a.length !== l); f = !0); } catch (r) { o = !0, n = r; } finally { try { if (!f && null != t.return && (u = t.return(), Object(u) !== u)) return; } finally { if (o) throw n; } } return a; } }
function _arrayWithHoles(r) { if (Array.isArray(r)) return r; }
function _classCallCheck(a, n) { if (!(a instanceof n)) throw new TypeError("Cannot call a class as a function"); }
function _defineProperties(e, r) { for (var t = 0; t < r.length; t++) { var o = r[t]; o.enumerable = o.enumerable || !1, o.configurable = !0, "value" in o && (o.writable = !0), Object.defineProperty(e, _toPropertyKey(o.key), o); } }
function _createClass(e, r, t) { return r && _defineProperties(e.prototype, r), t && _defineProperties(e, t), Object.defineProperty(e, "prototype", { writable: !1 }), e; }
function _defineProperty(e, r, t) { return (r = _toPropertyKey(r)) in e ? Object.defineProperty(e, r, { value: t, enumerable: !0, configurable: !0, writable: !0 }) : e[r] = t, e; }
function _toPropertyKey(t) { var i = _toPrimitive(t, "string"); return "symbol" == _typeof(i) ? i : i + ""; }
function _toPrimitive(t, r) { if ("object" != _typeof(t) || !t) return t; var e = t[Symbol.toPrimitive]; if (void 0 !== e) { var i = e.call(t, r || "default"); if ("object" != _typeof(i)) return i; throw new TypeError("@@toPrimitive must return a primitive value."); } return ("string" === r ? String : Number)(t); }
/**
* Supported providers
*/
var SUPPORTED_PROVIDERS = ['openai', 'groq'];
/**
* Validate provider
*
* @param provider
* @throws Error if provider is not supported
*/
var _validateProvider = function _validateProvider(provider) {
if (!SUPPORTED_PROVIDERS.includes(provider)) {
throw new Error("Unsupported provider \"".concat(provider, "\". Must be one of ").concat(SUPPORTED_PROVIDERS.join(', ')));
}
};
/**
* Zero-shot classifier
*/
var ZeroShotClassifier = /*#__PURE__*/function () {
function ZeroShotClassifier(config) {
_classCallCheck(this, ZeroShotClassifier);
/** Provider used for API */
_defineProperty(this, "provider", void 0);
/** Model used for classification */
_defineProperty(this, "model", void 0);
/** Provider API Key */
_defineProperty(this, "apiKey", void 0);
/** Labels to classify against */
_defineProperty(this, "labels", void 0);
/** Dimensions used for embeddings */
_defineProperty(this, "dimensions", void 0);
/** Labels cache */
_defineProperty(this, "labelsCache", void 0);
/** Data cache */
_defineProperty(this, "dataCache", void 0);
/** API client */
_defineProperty(this, "client", void 0);
var _config$model = config.model,
model = _config$model === void 0 ? 'text-embedding-3-small' : _config$model,
_config$provider = config.provider,
provider = _config$provider === void 0 ? 'openai' : _config$provider,
apiKey = config.apiKey,
_config$labels = config.labels,
labels = _config$labels === void 0 ? [] : _config$labels,
dimensions = config.dimensions,
_config$labelsCache = config.labelsCache,
labelsCache = _config$labelsCache === void 0 ? {} : _config$labelsCache,
_config$dataCache = config.dataCache,
dataCache = _config$dataCache === void 0 ? {} : _config$dataCache;
_validateProvider(provider);
this.model = model;
this.provider = provider;
this.apiKey = apiKey;
this.labels = labels || [];
this.dimensions = dimensions;
this.labelsCache = labelsCache || {};
this.dataCache = dataCache || {};
this._createAndSetClient({
model: model,
provider: provider,
apiKey: apiKey
});
}
/**
* Create and set client based on provider, model and config
* @param config
*/
return _createClass(ZeroShotClassifier, [{
key: "_createAndSetClient",
value: function _createAndSetClient(config) {
this.client = (0, _providers.getProvider)(config.provider).createClient(config);
}
/**
* Set labels to classify against and clear labels cache
* @param labels string[]
*/
}, {
key: "setLabels",
value: function setLabels(labels) {
this.labels = labels;
// Clear cache from stale entries based on updated labels
this.labelsCache = Object.fromEntries(Object.entries(this.labelsCache).filter(function (_ref) {
var _ref2 = _slicedToArray(_ref, 1),
key = _ref2[0];
return labels.includes(key);
}));
}
/**
* Set model and provider used for classification
* @param model string
* @param provider string
*/
}, {
key: "setProviderAndModel",
value: function setProviderAndModel() {
var provider = arguments.length > 0 && arguments[0] !== undefined ? arguments[0] : 'openai';
var model = arguments.length > 1 ? arguments[1] : undefined;
var apiKey = arguments.length > 2 ? arguments[2] : undefined;
var dimensions = arguments.length > 3 ? arguments[3] : undefined;
_validateProvider(provider);
// clear cache if provider or model changes
if (provider !== this.provider || model !== this.model || dimensions !== this.dimensions) {
this.clearAllCaches();
}
this.provider = provider;
this.model = model;
this.apiKey = apiKey;
this.dimensions = dimensions;
this._createAndSetClient({
model: this.model,
provider: this.provider,
apiKey: this.apiKey
});
}
/**
* Get embeddings in batches with concurrency
*/
}, {
key: "getEmbeddings",
value: (function () {
var _getEmbeddings = _asyncToGenerator(/*#__PURE__*/_regeneratorRuntime().mark(function _callee2(texts, batchSize, concurrency) {
var _this = this;
var type,
cache,
uncachedTexts,
chunks,
_args2 = arguments;
return _regeneratorRuntime().wrap(function _callee2$(_context2) {
while (1) switch (_context2.prev = _context2.next) {
case 0:
type = _args2.length > 3 && _args2[3] !== undefined ? _args2[3] : 'data';
cache = type === 'label' ? this.labelsCache : this.dataCache;
uncachedTexts = texts.filter(function (text) {
return !cache[text];
});
chunks = (0, _lodash.default)(uncachedTexts, batchSize);
_context2.next = 6;
return (0, _pMap.default)(chunks, /*#__PURE__*/function () {
var _ref3 = _asyncToGenerator(/*#__PURE__*/_regeneratorRuntime().mark(function _callee(currChunk) {
var response;
return _regeneratorRuntime().wrap(function _callee$(_context) {
while (1) switch (_context.prev = _context.next) {
case 0:
_context.next = 2;
return (0, _providers.getProvider)(_this.provider).createEmbedding(_this.client, {
model: _this.model,
input: currChunk,
dimensions: _this.dimensions
});
case 2:
response = _context.sent;
currChunk.forEach(function (text, index) {
cache[text] = response[index].embedding;
});
return _context.abrupt("return", response.map(function (r) {
return r.embedding;
}));
case 5:
case "end":
return _context.stop();
}
}, _callee);
}));
return function (_x4) {
return _ref3.apply(this, arguments);
};
}(), {
concurrency: concurrency
});
case 6:
return _context2.abrupt("return", texts.map(function (text) {
return cache[text];
}));
case 7:
case "end":
return _context2.stop();
}
}, _callee2, this);
}));
function getEmbeddings(_x, _x2, _x3) {
return _getEmbeddings.apply(this, arguments);
}
return getEmbeddings;
}()
/**
* Clear cache for data embeddings
*/
)
}, {
key: "clearDataCache",
value: function clearDataCache() {
this.dataCache = {};
}
/**
* Clear cache for label embeddings
*/
}, {
key: "clearLabelsCache",
value: function clearLabelsCache() {
this.labelsCache = {};
}
/**
* Clear all caches
*/
}, {
key: "clearAllCaches",
value: function clearAllCaches() {
this.clearDataCache();
this.clearLabelsCache();
}
/**
* Classify data
*
* @param data string[]
* @param config object
* @returns Object[]
*/
}, {
key: "classify",
value: (function () {
var _classify = _asyncToGenerator(/*#__PURE__*/_regeneratorRuntime().mark(function _callee5(data) {
var _this2 = this;
var config,
_config$similarity,
similarity,
_config$embeddingBatc,
embeddingBatchSizeData,
_config$embeddingBatc2,
embeddingBatchSizeLabels,
_config$embeddingConc,
embeddingConcurrencyData,
_config$embeddingConc2,
embeddingConcurrencyLabels,
_config$comparingConc,
comparingConcurrencyTop,
_config$comparingConc2,
comparingConcurrencyBottom,
_yield$Promise$all,
_yield$Promise$all2,
labelsEmbeddings,
dataEmbeddings,
getSimilarity,
result,
_args5 = arguments;
return _regeneratorRuntime().wrap(function _callee5$(_context5) {
while (1) switch (_context5.prev = _context5.next) {
case 0:
config = _args5.length > 1 && _args5[1] !== undefined ? _args5[1] : {};
if (this.labels.length) {
_context5.next = 3;
break;
}
throw new Error('Labels must be set.');
case 3:
_config$similarity = config.similarity, similarity = _config$similarity === void 0 ? 'cosine' : _config$similarity, _config$embeddingBatc = config.embeddingBatchSizeData, embeddingBatchSizeData = _config$embeddingBatc === void 0 ? 50 : _config$embeddingBatc, _config$embeddingBatc2 = config.embeddingBatchSizeLabels, embeddingBatchSizeLabels = _config$embeddingBatc2 === void 0 ? 50 : _config$embeddingBatc2, _config$embeddingConc = config.embeddingConcurrencyData, embeddingConcurrencyData = _config$embeddingConc === void 0 ? 5 : _config$embeddingConc, _config$embeddingConc2 = config.embeddingConcurrencyLabels, embeddingConcurrencyLabels = _config$embeddingConc2 === void 0 ? 5 : _config$embeddingConc2, _config$comparingConc = config.comparingConcurrencyTop, comparingConcurrencyTop = _config$comparingConc === void 0 ? 10 : _config$comparingConc, _config$comparingConc2 = config.comparingConcurrencyBottom, comparingConcurrencyBottom = _config$comparingConc2 === void 0 ? 10 : _config$comparingConc2; // Parallel embedding computation for labels and data
_context5.next = 6;
return Promise.all([this.getEmbeddings(this.labels, embeddingBatchSizeLabels, embeddingConcurrencyLabels, 'label'), this.getEmbeddings(data, embeddingBatchSizeData, embeddingConcurrencyData, 'data')]);
case 6:
_yield$Promise$all = _context5.sent;
_yield$Promise$all2 = _slicedToArray(_yield$Promise$all, 2);
labelsEmbeddings = _yield$Promise$all2[0];
dataEmbeddings = _yield$Promise$all2[1];
/** similarity getter function */
getSimilarity = (0, _similarityFunctions.getSimilarityFunction)(similarity);
_context5.next = 13;
return (0, _pMap.default)(dataEmbeddings, /*#__PURE__*/function () {
var _ref4 = _asyncToGenerator(/*#__PURE__*/_regeneratorRuntime().mark(function _callee4(dataEmbedding) {
var similarities, bestIndex;
return _regeneratorRuntime().wrap(function _callee4$(_context4) {
while (1) switch (_context4.prev = _context4.next) {
case 0:
_context4.next = 2;
return (0, _pMap.default)(labelsEmbeddings, /*#__PURE__*/function () {
var _ref5 = _asyncToGenerator(/*#__PURE__*/_regeneratorRuntime().mark(function _callee3(labelEmbedding) {
return _regeneratorRuntime().wrap(function _callee3$(_context3) {
while (1) switch (_context3.prev = _context3.next) {
case 0:
return _context3.abrupt("return", getSimilarity(dataEmbedding, labelEmbedding));
case 1:
case "end":
return _context3.stop();
}
}, _callee3);
}));
return function (_x7) {
return _ref5.apply(this, arguments);
};
}(), {
concurrency: comparingConcurrencyBottom
});
case 2:
similarities = _context4.sent;
// find closest label based on similarity
bestIndex = similarities.indexOf(Math[similarity === 'euclidean' ? 'min' : 'max'].apply(Math, _toConsumableArray(similarities)));
return _context4.abrupt("return", {
label: _this2.labels[bestIndex],
confidence: similarities[bestIndex] // Include confidence score
});
case 5:
case "end":
return _context4.stop();
}
}, _callee4);
}));
return function (_x6) {
return _ref4.apply(this, arguments);
};
}(), {
concurrency: comparingConcurrencyTop
});
case 13:
result = _context5.sent;
return _context5.abrupt("return", result);
case 15:
case "end":
return _context5.stop();
}
}, _callee5, this);
}));
function classify(_x5) {
return _classify.apply(this, arguments);
}
return classify;
}())
}]);
}();
var _default = exports.default = ZeroShotClassifier;
//# sourceMappingURL=index.js.map