Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion docs/content.zh/docs/core-concept/ai-model.md
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,17 @@ AI 模型可用于 transform 表达式中的文本生成、文本分析、embedd

## AI Functions

模型名称必须是字符串常量,并引用 `pipeline.model` 中声明的模型。文本、embedding 和图片函数分别要求模型客户端实现对应的 capability;Pipeline 会在执行前校验引用模型的 capability 是否匹配。
模型参数可以是任意 `STRING` 表达式,并会针对每条记录求值。因此可以使用字段、`IF` 或 `CASE` 动态选择模型。仅在实际调用 AI 函数时,才会根据求值结果查找 `pipeline.model` 中声明的模型;如果选中的模型未声明,或没有实现函数所需的 capability,当前记录会在运行时报错。

例如,下面的表达式会根据每条记录的优先级选择模型:

```sql
AI_COMPLETE(
IF(priority = 'high', 'powerful_model', 'economical_model'),
content,
'总结输入内容'
)
```

所有文本函数都会将模型返回的 JSON 解析为 `VARIANT`。

Expand Down
12 changes: 11 additions & 1 deletion docs/content/docs/core-concept/ai-model.md
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,17 @@ image understanding.

## AI Functions

The model name must be a string constant that refers to a model declared in `pipeline.model`. Text functions require a model client that implements text generation, while embedding and image functions require their corresponding capabilities. The pipeline validates the referenced model capability before execution.
The model argument accepts any `STRING` expression and is evaluated for each record. This enables dynamic model selection with a column, `IF`, or `CASE`. The selected name is resolved against the models declared in `pipeline.model` only when the AI function is invoked. If the selected model is undeclared or does not provide the capability required by the function, processing of that record fails at runtime.

For example, the following expression chooses a model based on each record's priority:

```sql
AI_COMPLETE(
IF(priority = 'high', 'powerful_model', 'economical_model'),
content,
'Summarize the input'
)
```

All text functions return `VARIANT` values parsed from the model's JSON response.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,17 +35,14 @@
import org.apache.flink.cdc.runtime.operators.transform.PostTransformOperatorBuilder;
import org.apache.flink.cdc.runtime.operators.transform.PreTransformOperator;
import org.apache.flink.cdc.runtime.operators.transform.PreTransformOperatorBuilder;
import org.apache.flink.cdc.runtime.parser.TransformParser;
import org.apache.flink.cdc.runtime.typeutils.EventTypeInfo;
import org.apache.flink.streaming.api.datastream.DataStream;
import org.apache.flink.streaming.api.environment.StreamExecutionEnvironment;

import java.util.Collections;
import java.util.HashSet;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.stream.Collectors;

/**
Expand All @@ -67,8 +64,6 @@ public DataStream<Event> translatePreTransform(
if (transforms.isEmpty()) {
return input;
}
validateModelReferences(
transforms, models, getUserDefinedFunctionNames(udfFunctions, models));
return input.transform(
"Transform:Schema",
new EventTypeInfo(),
Expand Down Expand Up @@ -147,8 +142,6 @@ public DataStream<Event> translatePostTransform(
.map(this::modelToUDFTuple)
.collect(Collectors.toList()));
Map<String, AiModelClient> modelClients = loadModelClients(models, env);
validateModelCapabilities(
transforms, modelClients, getUserDefinedFunctionNames(udfFunctions, models));
postTransformFunctionBuilder.addModelClients(modelClients);
return input.transform(
"Transform:Data", new EventTypeInfo(), postTransformFunctionBuilder.build())
Expand Down Expand Up @@ -191,48 +184,6 @@ private Map<String, AiModelClient> loadModelClients(
return clients;
}

private void validateModelReferences(
List<TransformDef> transforms,
List<ModelDef> models,
Set<String> userDefinedFunctionNames) {
Set<String> clientModelNames =
models.stream()
.filter(model -> !model.isLegacy())
.map(ModelDef::getName)
.collect(Collectors.toSet());
for (TransformDef transform : transforms) {
TransformParser.validateAiModelReferences(
transform.getProjection(),
transform.getFilter(),
clientModelNames,
userDefinedFunctionNames);
}
}

private void validateModelCapabilities(
List<TransformDef> transforms,
Map<String, AiModelClient> modelClients,
Set<String> userDefinedFunctionNames) {
for (TransformDef transform : transforms) {
TransformParser.validateAiModelCapabilities(
transform.getProjection(),
transform.getFilter(),
modelClients,
userDefinedFunctionNames);
}
}

private Set<String> getUserDefinedFunctionNames(
List<UdfDef> udfFunctions, List<ModelDef> models) {
Set<String> functionNames = new HashSet<>();
udfFunctions.stream().map(UdfDef::getName).forEach(functionNames::add);
models.stream()
.filter(ModelDef::isLegacy)
.map(ModelDef::getName)
.forEach(functionNames::add);
return functionNames;
}

private Tuple3<String, String, Map<String, String>> udfDefToUDFTuple(UdfDef udf) {
return Tuple3.of(udf.getName(), udf.getClasspath(), udf.getOptions());
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,20 @@ void testAiCompleteInProjection() throws Exception {
"Dummy model closed.");
}

@Test
void testDynamicModelSelectionInProjection() throws Exception {
String[] output =
runAiFunctionTest(
"id, content, "
+ "AI_COMPLETE(IF(id = 1, 'testModel', 'missingModel'), content, 'Complete the text') AS completed",
List.of(ModelDef.of("testModel", "dummy", Collections.emptyMap())));

assertThat(output)
.containsExactly(
"CreateTableEvent{tableId=default_namespace.default_schema.mytable1, schema=columns={`id` INT NOT NULL,`content` STRING,`completed` VARIANT}, primaryKeys=id, options=()}",
"DataChangeEvent{tableId=default_namespace.default_schema.mytable1, before=[], after=[1, I love this product, {\"result\":\"dummy response\"}], op=INSERT, meta=()}");
}

@Test
void testSpecializedTextAiFunctionsInProjection() throws Exception {
String[] output =
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@

import java.io.IOException;
import java.util.List;
import java.util.Map;

/** Built-in AI functions available to transform expressions. */
public class AiFunctions {
Expand All @@ -39,53 +40,89 @@ public class AiFunctions {

private AiFunctions() {}

public static BinaryVariant aiComplete(AiModelClient model, String input, String systemPrompt) {
return generateText(model, AiTextFunctionDef.AI_COMPLETE, input, systemPrompt);
public static BinaryVariant aiComplete(
String modelName,
String input,
String systemPrompt,
Map<String, AiModelClient> modelClients) {
return generateText(
modelClients, modelName, AiTextFunctionDef.AI_COMPLETE, input, systemPrompt);
}

public static BinaryVariant aiClassify(AiModelClient model, String input, String labels) {
return generateText(model, AiTextFunctionDef.AI_CLASSIFY, input, labels);
public static BinaryVariant aiClassify(
String modelName,
String input,
String labels,
Map<String, AiModelClient> modelClients) {
return generateText(modelClients, modelName, AiTextFunctionDef.AI_CLASSIFY, input, labels);
}

public static BinaryVariant aiTranslate(
AiModelClient model, String input, String sourceLang, String targetLang) {
return generateText(model, AiTextFunctionDef.AI_TRANSLATE, input, sourceLang, targetLang);
String modelName,
String input,
String sourceLang,
String targetLang,
Map<String, AiModelClient> modelClients) {
return generateText(
modelClients,
modelName,
AiTextFunctionDef.AI_TRANSLATE,
input,
sourceLang,
targetLang);
}

public static BinaryVariant aiSummarize(AiModelClient model, String input, int maxLength) {
return generateText(model, AiTextFunctionDef.AI_SUMMARIZE, input, maxLength);
public static BinaryVariant aiSummarize(
String modelName,
String input,
int maxLength,
Map<String, AiModelClient> modelClients) {
return generateText(
modelClients, modelName, AiTextFunctionDef.AI_SUMMARIZE, input, maxLength);
}

public static BinaryVariant aiSentiment(AiModelClient model, String input) {
return generateText(model, AiTextFunctionDef.AI_SENTIMENT, input);
public static BinaryVariant aiSentiment(
String modelName, String input, Map<String, AiModelClient> modelClients) {
return generateText(modelClients, modelName, AiTextFunctionDef.AI_SENTIMENT, input);
}

public static BinaryVariant aiExtract(AiModelClient model, String input, String schema) {
return generateText(model, AiTextFunctionDef.AI_EXTRACT, input, schema);
public static BinaryVariant aiExtract(
String modelName,
String input,
String schema,
Map<String, AiModelClient> modelClients) {
return generateText(modelClients, modelName, AiTextFunctionDef.AI_EXTRACT, input, schema);
}

public static BinaryVariant aiMask(AiModelClient model, String input, String entities) {
return generateText(model, AiTextFunctionDef.AI_MASK, input, entities);
public static BinaryVariant aiMask(
String modelName,
String input,
String entities,
Map<String, AiModelClient> modelClients) {
return generateText(modelClients, modelName, AiTextFunctionDef.AI_MASK, input, entities);
}

private static BinaryVariant generateText(
AiModelClient model,
Map<String, AiModelClient> modelClients,
String modelName,
AiTextFunctionDef function,
String input,
Object... promptArguments) {
if (input == null) {
return null;
}
if (!(model instanceof SupportsTextGeneration)) {
throw new UnsupportedOperationException(
"Model " + model.getClass().getName() + " does not support text generation");
}
SupportsTextGeneration model =
resolveModel(
modelClients,
modelName,
function.getFunctionName(),
SupportsTextGeneration.class);

String prompt =
function.buildPrompt(promptArguments)
+ "\n"
+ buildOutputSchemaHint(function.getOutputType());
String json = ((SupportsTextGeneration) model).generate(prompt, input);
String json = model.generate(prompt, input);
if (json == null) {
return null;
}
Expand All @@ -101,43 +138,77 @@ private static BinaryVariant generateText(
}
}

public static List<Float> aiEmbed(AiModelClient model, String input) {
public static List<Float> aiEmbed(
String modelName, String input, Map<String, AiModelClient> modelClients) {
if (input == null) {
return null;
}
if (!(model instanceof SupportsEmbedding)) {
throw new UnsupportedOperationException(
"Model " + model.getClass().getName() + " does not support embedding");
}
float[] embedding = ((SupportsEmbedding) model).embed(input);
SupportsEmbedding model =
resolveModel(modelClients, modelName, "AI_EMBED", SupportsEmbedding.class);
float[] embedding = model.embed(input);
return embedding == null ? null : Floats.asList(embedding);
}

/** Dispatches image-to-text AI functions. */
public static String aiImageComplete(AiModelClient model, byte[] image, String prompt) {
public static String aiImageComplete(
String modelName,
byte[] image,
String prompt,
Map<String, AiModelClient> modelClients) {
if (image == null) {
return null;
}
if (!(model instanceof SupportsImageTextGeneration)) {
throw new UnsupportedOperationException(
"Model "
+ model.getClass().getName()
+ " does not support image text generation");
}
return ((SupportsImageTextGeneration) model).generateTextFromImage(image, prompt);
SupportsImageTextGeneration model =
resolveModel(
modelClients,
modelName,
"AI_IMAGE_COMPLETE",
SupportsImageTextGeneration.class);
return model.generateTextFromImage(image, prompt);
}

/** Dispatches image embedding AI functions. */
public static List<Float> aiImageEmbed(AiModelClient model, byte[] image) {
public static List<Float> aiImageEmbed(
String modelName, byte[] image, Map<String, AiModelClient> modelClients) {
if (image == null) {
return null;
}
if (!(model instanceof SupportsImageEmbedding)) {
SupportsImageEmbedding model =
resolveModel(
modelClients, modelName, "AI_IMAGE_EMBED", SupportsImageEmbedding.class);
float[] embedding = model.embedImage(image);
return embedding == null ? null : Floats.asList(embedding);
}

private static <T> T resolveModel(
Map<String, AiModelClient> modelClients,
String modelName,
String functionName,
Class<T> requiredCapability) {
if (modelName == null) {
throw new IllegalArgumentException(
"Model name referenced by " + functionName + " must not be null.");
}
AiModelClient model = modelClients.get(modelName);
if (model == null) {
throw new IllegalArgumentException(
"Model '"
+ modelName
+ "' referenced by "
+ functionName
+ " has not been declared.");
}
if (!requiredCapability.isInstance(model)) {
throw new UnsupportedOperationException(
"Model " + model.getClass().getName() + " does not support image embedding");
"Model '"
+ modelName
+ "' could not be used in "
+ functionName
+ " because it does not implement "
+ requiredCapability.getSimpleName()
+ " interface.");
}
float[] embedding = ((SupportsImageEmbedding) model).embedImage(image);
return embedding == null ? null : Floats.asList(embedding);
return requiredCapability.cast(model);
}

private static String truncateInvalidJsonResponse(String response) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ public ProjectionColumnProcessor(
this.transformExpressionKey = generateTransformExpressionKey();
this.expressionEvaluator =
TransformExpressionCompiler.compileExpression(
transformExpressionKey, udfDescriptors, modelClients);
transformExpressionKey, udfDescriptors);
this.udfFunctionInstances = udfFunctionInstances;
}

Expand Down Expand Up @@ -148,8 +148,8 @@ private Object[] generateParams(Object[] rowData, TransformContext context) {
// 3 - Add UDF function instances
params.addAll(udfFunctionInstances);

// 4 - Add AI model client instances
params.addAll(modelClients.values());
// 4 - Add AI model clients
params.add(modelClients);
return params.toArray();
}

Expand Down
Loading
Loading