feat: support transform component callable (#1072)

This commit is contained in:
Alex Yang
2024-07-25 11:19:57 -07:00
committed by GitHub
parent 086b94038e
commit 91d02a4fc0
14 changed files with 168 additions and 125 deletions
+7
View File
@@ -0,0 +1,7 @@
---
"@llamaindex/core": patch
"llamaindex": patch
"@llamaindex/core-e2e": patch
---
feat: support transform component callable
+21 -18
View File
@@ -1,7 +1,6 @@
import { type Tokenizers } from "@llamaindex/env";
import type { MessageContentDetail } from "../llms";
import type { TransformComponent } from "../schema";
import { BaseNode, MetadataMode } from "../schema";
import { BaseNode, MetadataMode, TransformComponent } from "../schema";
import { extractSingleText } from "../utils";
import { truncateMaxTokens } from "./tokenizer.js";
import { SimilarityType, similarity } from "./utils.js";
@@ -20,10 +19,29 @@ export type BaseEmbeddingOptions = {
logProgress?: boolean;
};
export abstract class BaseEmbedding implements TransformComponent {
export abstract class BaseEmbedding extends TransformComponent {
embedBatchSize = DEFAULT_EMBED_BATCH_SIZE;
embedInfo?: EmbeddingInfo;
constructor() {
super(
async (
nodes: BaseNode[],
options?: BaseEmbeddingOptions,
): Promise<BaseNode[]> => {
const texts = nodes.map((node) => node.getContent(MetadataMode.EMBED));
const embeddings = await this.getTextEmbeddingsBatch(texts, options);
for (let i = 0; i < nodes.length; i++) {
nodes[i].embedding = embeddings[i];
}
return nodes;
},
);
}
similarity(
embedding1: number[],
embedding2: number[],
@@ -76,21 +94,6 @@ export abstract class BaseEmbedding implements TransformComponent {
);
}
async transform(
nodes: BaseNode[],
options?: BaseEmbeddingOptions,
): Promise<BaseNode[]> {
const texts = nodes.map((node) => node.getContent(MetadataMode.EMBED));
const embeddings = await this.getTextEmbeddingsBatch(texts, options);
for (let i = 0; i < nodes.length; i++) {
nodes[i].embedding = embeddings[i];
}
return nodes;
}
truncateMaxTokens(input: string[]): string[] {
return input.map((s) => {
// truncate to max tokens
+8 -6
View File
@@ -5,13 +5,19 @@ import {
MetadataMode,
NodeRelationship,
TextNode,
type TransformComponent,
TransformComponent,
} from "../schema";
export abstract class NodeParser implements TransformComponent {
export abstract class NodeParser extends TransformComponent {
includeMetadata: boolean = true;
includePrevNextRel: boolean = true;
constructor() {
super(async (nodes: BaseNode[]): Promise<BaseNode[]> => {
return this.getNodesFromDocuments(nodes as TextNode[]);
});
}
protected postProcessParsedNodes(
nodes: TextNode[],
parentDocMap: Map<string, TextNode>,
@@ -90,10 +96,6 @@ export abstract class NodeParser implements TransformComponent {
return nodes;
}
async transform(nodes: BaseNode[], options?: {}): Promise<BaseNode[]> {
return this.getNodesFromDocuments(nodes as TextNode[]);
}
}
export abstract class TextSplitter extends NodeParser {
+1 -1
View File
@@ -1,4 +1,4 @@
export * from "./node";
export type { TransformComponent } from "./type";
export { TransformComponent } from "./type";
export { EngineResponse } from "./type/engineresponse";
export * from "./zod";
+24 -2
View File
@@ -1,8 +1,30 @@
import { randomUUID } from "@llamaindex/env";
import type { BaseNode } from "./node";
export interface TransformComponent {
transform<Options extends Record<string, unknown>>(
interface TransformComponentSignature {
<Options extends Record<string, unknown>>(
nodes: BaseNode[],
options?: Options,
): Promise<BaseNode[]>;
}
export interface TransformComponent extends TransformComponentSignature {
id: string;
}
export class TransformComponent {
constructor(transformFn: TransformComponentSignature) {
Object.defineProperties(
transformFn,
Object.getOwnPropertyDescriptors(this.constructor.prototype),
);
const transform = function transform(
...args: Parameters<TransformComponentSignature>
) {
return transformFn(...args);
};
Reflect.setPrototypeOf(transform, new.target.prototype);
transform.id = randomUUID();
return transform;
}
}
@@ -1,15 +1,26 @@
import { TransformComponent } from "@llamaindex/core/schema";
import {
BaseEmbedding,
BaseNode,
SimilarityType,
type BaseEmbedding,
type EmbeddingInfo,
type MessageContentDetail,
} from "llamaindex";
export class OpenAIEmbedding implements BaseEmbedding {
export class OpenAIEmbedding
extends TransformComponent
implements BaseEmbedding
{
embedInfo?: EmbeddingInfo | undefined;
embedBatchSize = 512;
constructor() {
super(async (nodes: BaseNode[], _options?: any): Promise<BaseNode[]> => {
nodes.forEach((node) => (node.embedding = [0]));
return nodes;
});
}
async getQueryEmbedding(query: MessageContentDetail) {
return [0];
}
@@ -34,11 +45,6 @@ export class OpenAIEmbedding implements BaseEmbedding {
return 1;
}
async transform(nodes: BaseNode[], _options?: any): Promise<BaseNode[]> {
nodes.forEach((node) => (node.embedding = [0]));
return nodes;
}
truncateMaxTokens(input: string[]): string[] {
return input;
}
+17 -11
View File
@@ -1,11 +1,15 @@
import type { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import { MetadataMode, TextNode } from "@llamaindex/core/schema";
import {
BaseNode,
MetadataMode,
TextNode,
TransformComponent,
} from "@llamaindex/core/schema";
import { defaultNodeTextTemplate } from "./prompts.js";
/*
* Abstract class for all extractors.
*/
export abstract class BaseExtractor implements TransformComponent {
export abstract class BaseExtractor extends TransformComponent {
isTextNodeOnly: boolean = true;
showProgress: boolean = true;
metadataMode: MetadataMode = MetadataMode.ALL;
@@ -13,16 +17,18 @@ export abstract class BaseExtractor implements TransformComponent {
inPlace: boolean = true;
numWorkers: number = 4;
abstract extract(nodes: BaseNode[]): Promise<Record<string, any>[]>;
async transform(nodes: BaseNode[], options?: any): Promise<BaseNode[]> {
return this.processNodes(
nodes,
options?.excludedEmbedMetadataKeys,
options?.excludedLlmMetadataKeys,
);
constructor() {
super(async (nodes: BaseNode[], options?: any): Promise<BaseNode[]> => {
return this.processNodes(
nodes,
options?.excludedEmbedMetadataKeys,
options?.excludedLlmMetadataKeys,
);
});
}
abstract extract(nodes: BaseNode[]): Promise<Record<string, any>[]>;
/**
*
* @param nodes Nodes to extract metadata from.
@@ -172,7 +172,7 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
const embedModel =
this.embedModel ?? this.vectorStores[type as ModalityType]?.embedModel;
if (embedModel && nodes) {
await embedModel.transform(nodes, {
await embedModel(nodes, {
logProgress: options?.logProgress,
});
}
@@ -35,7 +35,7 @@ export function getTransformationHash(
const transformString: string = transformToJSON(transform);
const hash = createSHA256();
hash.update(nodesStr + transformString);
hash.update(nodesStr + transformString + transform.id);
return hash.digest();
}
@@ -40,7 +40,7 @@ export async function runTransformations(
nodes = [...nodesToRun];
}
if (docStoreStrategy) {
nodes = await docStoreStrategy.transform(nodes);
nodes = await docStoreStrategy(nodes);
}
for (const transform of transformations) {
if (cache) {
@@ -49,11 +49,11 @@ export async function runTransformations(
if (cachedNodes) {
nodes = cachedNodes;
} else {
nodes = await transform.transform(nodes, transformOptions);
nodes = await transform(nodes, transformOptions);
await cache.put(hash, nodes);
}
} else {
nodes = await transform.transform(nodes, transformOptions);
nodes = await transform(nodes, transformOptions);
}
}
return nodes;
@@ -1,31 +1,30 @@
import type { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import type { BaseDocumentStore } from "../../storage/docStore/types.js";
/**
* Handle doc store duplicates by checking all hashes.
*/
export class DuplicatesStrategy implements TransformComponent {
export class DuplicatesStrategy extends TransformComponent {
private docStore: BaseDocumentStore;
constructor(docStore: BaseDocumentStore) {
super(async (nodes: BaseNode[]): Promise<BaseNode[]> => {
const hashes = await this.docStore.getAllDocumentHashes();
const currentHashes = new Set<string>();
const nodesToRun: BaseNode[] = [];
for (const node of nodes) {
if (!(node.hash in hashes) && !currentHashes.has(node.hash)) {
await this.docStore.setDocumentHash(node.id_, node.hash);
nodesToRun.push(node);
currentHashes.add(node.hash);
}
}
await this.docStore.addDocuments(nodesToRun, true);
return nodesToRun;
});
this.docStore = docStore;
}
async transform(nodes: BaseNode[]): Promise<BaseNode[]> {
const hashes = await this.docStore.getAllDocumentHashes();
const currentHashes = new Set<string>();
const nodesToRun: BaseNode[] = [];
for (const node of nodes) {
if (!(node.hash in hashes) && !currentHashes.has(node.hash)) {
await this.docStore.setDocumentHash(node.id_, node.hash);
nodesToRun.push(node);
currentHashes.add(node.hash);
}
}
await this.docStore.addDocuments(nodesToRun, true);
return nodesToRun;
}
}
@@ -1,4 +1,4 @@
import type { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import type { BaseDocumentStore } from "../../storage/docStore/types.js";
import type { VectorStore } from "../../storage/vectorStore/types.js";
import { classify } from "./classify.js";
@@ -7,43 +7,42 @@ import { classify } from "./classify.js";
* Handle docstore upserts by checking hashes and ids.
* Identify missing docs and delete them from docstore and vector store
*/
export class UpsertsAndDeleteStrategy implements TransformComponent {
export class UpsertsAndDeleteStrategy extends TransformComponent {
protected docStore: BaseDocumentStore;
protected vectorStores?: VectorStore[];
constructor(docStore: BaseDocumentStore, vectorStores?: VectorStore[]) {
super(async (nodes: BaseNode[]): Promise<BaseNode[]> => {
const { dedupedNodes, missingDocs, unusedDocs } = await classify(
this.docStore,
nodes,
);
// remove unused docs
for (const refDocId of unusedDocs) {
await this.docStore.deleteRefDoc(refDocId, false);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(refDocId);
}
}
}
// remove missing docs
for (const docId of missingDocs) {
await this.docStore.deleteDocument(docId, true);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(docId);
}
}
}
await this.docStore.addDocuments(dedupedNodes, true);
return dedupedNodes;
});
this.docStore = docStore;
this.vectorStores = vectorStores;
}
async transform(nodes: BaseNode[]): Promise<BaseNode[]> {
const { dedupedNodes, missingDocs, unusedDocs } = await classify(
this.docStore,
nodes,
);
// remove unused docs
for (const refDocId of unusedDocs) {
await this.docStore.deleteRefDoc(refDocId, false);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(refDocId);
}
}
}
// remove missing docs
for (const docId of missingDocs) {
await this.docStore.deleteDocument(docId, true);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(docId);
}
}
}
await this.docStore.addDocuments(dedupedNodes, true);
return dedupedNodes;
}
}
@@ -1,4 +1,4 @@
import type { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import { BaseNode, TransformComponent } from "@llamaindex/core/schema";
import type { BaseDocumentStore } from "../../storage/docStore/types.js";
import type { VectorStore } from "../../storage/vectorStore/types.js";
import { classify } from "./classify.js";
@@ -6,28 +6,27 @@ import { classify } from "./classify.js";
/**
* Handles doc store upserts by checking hashes and ids.
*/
export class UpsertsStrategy implements TransformComponent {
export class UpsertsStrategy extends TransformComponent {
protected docStore: BaseDocumentStore;
protected vectorStores?: VectorStore[];
constructor(docStore: BaseDocumentStore, vectorStores?: VectorStore[]) {
super(async (nodes: BaseNode[]): Promise<BaseNode[]> => {
const { dedupedNodes, unusedDocs } = await classify(this.docStore, nodes);
// remove unused docs
for (const refDocId of unusedDocs) {
await this.docStore.deleteRefDoc(refDocId, false);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(refDocId);
}
}
}
// add non-duplicate docs
await this.docStore.addDocuments(dedupedNodes, true);
return dedupedNodes;
});
this.docStore = docStore;
this.vectorStores = vectorStores;
}
async transform(nodes: BaseNode[]): Promise<BaseNode[]> {
const { dedupedNodes, unusedDocs } = await classify(this.docStore, nodes);
// remove unused docs
for (const refDocId of unusedDocs) {
await this.docStore.deleteRefDoc(refDocId, false);
if (this.vectorStores) {
for (const vectorStore of this.vectorStores) {
await vectorStore.delete(refDocId);
}
}
}
// add non-duplicate docs
await this.docStore.addDocuments(dedupedNodes, true);
return dedupedNodes;
}
}
@@ -1,4 +1,4 @@
import type { TransformComponent } from "@llamaindex/core/schema";
import { TransformComponent } from "@llamaindex/core/schema";
import type { BaseDocumentStore } from "../../storage/docStore/types.js";
import type { VectorStore } from "../../storage/vectorStore/types.js";
import { DuplicatesStrategy } from "./DuplicatesStrategy.js";
@@ -19,9 +19,9 @@ export enum DocStoreStrategy {
NONE = "none", // no-op strategy
}
class NoOpStrategy implements TransformComponent {
async transform(nodes: any[]): Promise<any[]> {
return nodes;
class NoOpStrategy extends TransformComponent {
constructor() {
super(async (nodes) => nodes);
}
}