This is an automated email from the ASF dual-hosted git repository.
davidzollo pushed a commit to branch dev
in repository https://gitbox.apache.org/repos/asf/seatunnel.git
The following commit(s) were added to refs/heads/dev by this push:
new 1863e33792 [Improve][Transform-V2][Embedding]Enhance multimodal
embeddings (#9996)
1863e33792 is described below
commit 1863e337921e1c9e93422d12b72180905d3da1c1
Author: loupipalien <[email protected]>
AuthorDate: Mon Jun 29 22:21:31 2026 +0800
[Improve][Transform-V2][Embedding]Enhance multimodal embeddings (#9996)
---
docs/en/transforms/embedding.md | 54 +++-
docs/zh/transforms/embedding.md | 54 +++-
.../resources/embedding_transform_multimodal.conf | 78 +++++
.../nlpmodel/embedding/EmbeddingTransform.java | 98 +++---
.../transform/nlpmodel/embedding/SrcField.java | 84 +++++
.../{FieldSpec.java => SrcFieldSpec.java} | 65 ++--
.../nlpmodel/embedding/VectorFieldSpec.java | 101 ++++++
.../embedding/multimodal/MultimodalFieldValue.java | 43 +--
.../embedding/remote/doubao/DoubaoModel.java | 64 ++--
.../embedding/DoubaoMultimodalModelTest.java | 355 +++++++++++++++------
.../transform/embedding/FieldSpecTest.java | 114 -------
.../transform/embedding/MultimodalConfigTest.java | 112 ++++++-
.../transform/embedding/VectorFieldSpecTest.java | 200 ++++++++++++
13 files changed, 1066 insertions(+), 356 deletions(-)
diff --git a/docs/en/transforms/embedding.md b/docs/en/transforms/embedding.md
index 98dfa0af01..6823332d4e 100644
--- a/docs/en/transforms/embedding.md
+++ b/docs/en/transforms/embedding.md
@@ -149,6 +149,58 @@ vectorization_fields {
}
```
+**Multi-field Mixing Multimodal Vectorization:**
+> Note: Currently, only the `DOUBAO` provider supports multimodal data
processing.
+```hocon
+vectorization_fields {
+ # Multi-field text
+ multi_field_text_vector = [product_name, description]
+
+ # Multi-field image
+ multi_field_image_vector = [
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = thumbnail_image
+ modality = png
+ format = url
+ }
+ ]
+
+ # Multi-field video
+ multi_field_video_vector = [
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ },
+ {
+ field = promotional_video
+ modality = mov
+ format = url
+ }
+ ]
+
+ # Multi-field mix multimodal
+ multi_field_mix_vector = [
+ product_name,
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ }
+ ]
+}
+```
+
**Field Specification Formats:**
**Supported Modality Types:**
@@ -162,7 +214,7 @@ vectorization_fields {
- `binary` - Binary data format
**Automatic Modality Detection:**
-When `modality` is not explicitly specified and `format` is not `binary`, the
system automatically detects the modality type based on the file suffix of the
field value:
+When `modality` is not explicitly specified and `format` is `url`, the system
automatically detects the modality type based on the file suffix of the field
value:
> **Important:** When using multimodal fields (image or video), ensure your
> model provider supports multimodal embedding. Image and video fields must
> contain valid URLs or binary data. Currently, `DOUBAO` provider supports
> multimodal data processing.
diff --git a/docs/zh/transforms/embedding.md b/docs/zh/transforms/embedding.md
index b8ace6ca6b..e80f96a575 100644
--- a/docs/zh/transforms/embedding.md
+++ b/docs/zh/transforms/embedding.md
@@ -134,6 +134,58 @@ vectorization_fields {
}
```
+**多字段混合多模态向量化:**
+> 注意: 目前,仅 `DOUBAO` 提供商支持多模态数据处理
+```hocon
+vectorization_fields {
+ # 多字段文本
+ multi_field_text_vector = [product_name, description]
+
+ # 多字段图片
+ multi_field_image_vector = [
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = thumbnail_image
+ modality = png
+ format = url
+ }
+ ]
+
+ # 多字段视频
+ multi_field_video_vector = [
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ },
+ {
+ field = promotional_video
+ modality = mov
+ format = url
+ }
+ ]
+
+ # 多字段混合多模态
+ multi_field_mix_vector = [
+ product_name,
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ }
+ ]
+}
+```
+
**字段规范格式:**
**支持的模态类型:**
@@ -147,7 +199,7 @@ vectorization_fields {
- `binary` - 二进制数据格式
**自动模态检测:**
-当未显式指定 `modality` 且 `format` 不是 `binary` 时,系统会根据字段值的文件后缀自动检测模态类型:
+当未显式指定 `modality` 且 `format` 是 `url` 时,系统会根据字段值的文件后缀自动检测模态类型:
> **重要:** 使用多模态字段(图片或视频)时,请确保您的模型提供商支持多模态 embedding。图片和视频字段必须包含有效的 URL
> 或二进制数据。目前,`DOUBAO` 提供商支持多模态数据处理。
diff --git
a/seatunnel-e2e/seatunnel-transforms-v2-e2e/seatunnel-transforms-v2-e2e-part-1/src/test/resources/embedding_transform_multimodal.conf
b/seatunnel-e2e/seatunnel-transforms-v2-e2e/seatunnel-transforms-v2-e2e-part-1/src/test/resources/embedding_transform_multimodal.conf
index efc7273145..cada8b7ca2 100644
---
a/seatunnel-e2e/seatunnel-transforms-v2-e2e/seatunnel-transforms-v2-e2e-part-1/src/test/resources/embedding_transform_multimodal.conf
+++
b/seatunnel-e2e/seatunnel-transforms-v2-e2e/seatunnel-transforms-v2-e2e-part-1/src/test/resources/embedding_transform_multimodal.conf
@@ -154,6 +154,48 @@ transform {
}
product_name_vector = product_name
+
+ multi_field_text_vector = [product_name, description]
+
+ multi_field_image_vector = [
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = thumbnail_image
+ modality = png
+ format = url
+ }
+ ]
+
+ multi_field_video_vector = [
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ },
+ {
+ field = promotional_video
+ modality = mov
+ format = url
+ }
+ ]
+
+ multi_field_mix_vector = [
+ product_name,
+ {
+ field = product_image_url
+ modality = jpeg
+ format = url
+ },
+ {
+ field = product_video_url
+ modality = mp4
+ format = url
+ }
+ ]
}
plugin_output = "multimodal_embedding_output"
@@ -219,6 +261,42 @@ sink {
}
]
},
+ {
+ field_name = multi_field_text_vector
+ field_type = float_vector
+ field_value = [
+ {
+ rule_type = NOT_NULL
+ }
+ ]
+ },
+ {
+ field_name = multi_field_image_vector
+ field_type = float_vector
+ field_value = [
+ {
+ rule_type = NOT_NULL
+ }
+ ]
+ },
+ {
+ field_name = multi_field_video_vector
+ field_type = float_vector
+ field_value = [
+ {
+ rule_type = NOT_NULL
+ }
+ ]
+ },
+ {
+ field_name = multi_field_mix_vector
+ field_type = float_vector
+ field_value = [
+ {
+ rule_type = NOT_NULL
+ }
+ ]
+ },
{
field_name = category
field_type = string
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/EmbeddingTransform.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/EmbeddingTransform.java
index 6e7ba72a3a..bcf09b67bb 100644
---
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/EmbeddingTransform.java
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/EmbeddingTransform.java
@@ -52,22 +52,22 @@ import java.io.IOException;
import java.net.URISyntaxException;
import java.nio.ByteBuffer;
import java.util.ArrayList;
-import java.util.HashMap;
+import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.TreeMap;
import java.util.concurrent.ConcurrentHashMap;
+import java.util.stream.Collectors;
@Slf4j
public class EmbeddingTransform extends MultipleFieldOutputTransform {
private final ReadonlyConfig config;
- private List<Integer> fieldOriginalIndexes;
private transient Model model;
private Integer dimension;
private boolean isMultimodalFields = false;
- private Map<Integer, FieldSpec> fieldSpecMap;
+ private Map<VectorFieldSpec, List<Integer>> fieldSpecMap;
private List<String> fieldNames;
private final Map<String, TreeMap<Long, byte[]>> binaryFileCache = new
ConcurrentHashMap<>();
@@ -204,30 +204,35 @@ public class EmbeddingTransform extends
MultipleFieldOutputTransform {
}
private void initOutputFields(SeaTunnelRowType inputRowType,
ReadonlyConfig config) {
- Map<Integer, FieldSpec> fieldSpecMap = new HashMap<>();
- List<String> fieldNames = new ArrayList<>();
Map<String, Object> fieldsConfig =
config.get(EmbeddingTransformConfig.VECTORIZATION_FIELDS);
if (fieldsConfig == null || fieldsConfig.isEmpty()) {
throw new IllegalArgumentException("vectorization_fields
configuration is required");
}
- for (Map.Entry<String, Object> field : fieldsConfig.entrySet()) {
- FieldSpec fieldSpec = new FieldSpec(field);
- log.info("Field spec: {}", fieldSpec.toString());
- String srcField = fieldSpec.getFieldName();
- int srcFieldIndex;
- try {
- srcFieldIndex = inputRowType.indexOf(srcField);
- } catch (IllegalArgumentException e) {
- throw
TransformCommonError.cannotFindInputFieldError(getPluginName(), srcField);
- }
- if (fieldSpec.isMultimodalField()) {
- isMultimodalFields = true;
+ List<String> fieldNames = new ArrayList<>();
+ Map<VectorFieldSpec, List<Integer>> fieldSpecMap = new
LinkedHashMap<>();
+ for (Map.Entry<String, Object> fieldConfig : fieldsConfig.entrySet()) {
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(fieldConfig);
+ log.info("Vector field spec: {}", vectorFieldSpec);
+ List<String> srcFieldNames =
+ vectorFieldSpec.getSrcFieldSpecs().stream()
+ .map(SrcFieldSpec::getFieldName)
+ .collect(Collectors.toList());
+ List<Integer> srcFieldIndexes = new ArrayList<>();
+ for (String srcFieldName : srcFieldNames) {
+ try {
+ srcFieldIndexes.add(inputRowType.indexOf(srcFieldName));
+ } catch (IllegalArgumentException e) {
+ throw TransformCommonError.cannotFindInputFieldsError(
+ getPluginName(), srcFieldNames);
+ }
}
- fieldSpecMap.put(srcFieldIndex, fieldSpec);
- fieldNames.add(field.getKey());
+ fieldSpecMap.put(vectorFieldSpec, srcFieldIndexes);
+ fieldNames.add(vectorFieldSpec.getFieldName());
}
+ this.isMultimodalFields =
+
fieldSpecMap.keySet().stream().anyMatch(VectorFieldSpec::isMultimodalField);
this.fieldSpecMap = fieldSpecMap;
this.fieldNames = fieldNames;
}
@@ -239,19 +244,28 @@ public class EmbeddingTransform extends
MultipleFieldOutputTransform {
if (MetadataUtil.isBinaryFormat(inputRow)) {
return vectorizationBinaryRow(inputRow);
}
- Set<Integer> fieldOriginalIndexes = fieldSpecMap.keySet();
- Object[] fieldValues = new Object[fieldOriginalIndexes.size()];
- List<ByteBuffer> vectorization;
+
+ Set<VectorFieldSpec> vectorFieldSpecs = fieldSpecMap.keySet();
+ Object[] fieldValues = new Object[vectorFieldSpecs.size()];
int i = 0;
- for (Integer fieldOriginalIndex : fieldOriginalIndexes) {
- FieldSpec fieldSpec = fieldSpecMap.get(fieldOriginalIndex);
- Object value = inputRow.getField(fieldOriginalIndex);
+ for (VectorFieldSpec vectorFieldSpec : vectorFieldSpecs) {
+ List<SrcFieldSpec> srcFieldSpecs =
vectorFieldSpec.getSrcFieldSpecs();
+ List<Integer> srcFieldIndexes =
fieldSpecMap.get(vectorFieldSpec);
+ List<SrcField> srcFields = new ArrayList<>();
+ for (int j = 0; j < srcFieldSpecs.size(); j++) {
+ srcFields.add(
+ new SrcField(
+ srcFieldSpecs.get(j),
+
inputRow.getField(srcFieldIndexes.get(j))));
+ }
fieldValues[i++] =
- isMultimodalFields ? new
MultimodalFieldValue(fieldSpec, value) : value;
+ isMultimodalFields
+ ? new MultimodalFieldValue(srcFields)
+ : srcFields.get(0).getFieldValue();
}
- vectorization = model.vectorization(fieldValues);
+ List<ByteBuffer> vectorization = model.vectorization(fieldValues);
return vectorization.toArray();
} catch (Exception e) {
throw new RuntimeException("Failed to data vectorization", e);
@@ -289,32 +303,34 @@ public class EmbeddingTransform extends
MultipleFieldOutputTransform {
/** Process a row in binary format: [data, relativePath, partIndex] */
private Object[] vectorizationBinaryRow(SeaTunnelRowAccessor inputRow)
throws Exception {
-
byte[] completeData = processBinaryRow(inputRow);
if (completeData == null) {
return null;
}
- Set<Integer> fieldOriginalIndexes = fieldSpecMap.keySet();
- Object[] fieldValues = new Object[fieldOriginalIndexes.size()];
+
+ Set<VectorFieldSpec> vectorFieldSpecs = fieldSpecMap.keySet();
+ Object[] fieldValues = new Object[vectorFieldSpecs.size()];
int i = 0;
- for (Integer fieldOriginalIndex : fieldOriginalIndexes) {
- FieldSpec fieldSpec = fieldSpecMap.get(fieldOriginalIndex);
- if (fieldSpec.isBinary()) {
- fieldValues[i++] = new MultimodalFieldValue(fieldSpec,
completeData);
- } else {
- log.warn(
- "Non-binary field {} configured in binary format data",
- fieldSpec.getFieldName());
- fieldValues[i++] = null;
+ for (VectorFieldSpec vectorFieldSpec : vectorFieldSpecs) {
+ List<SrcFieldSpec> srcFieldSpecs =
vectorFieldSpec.getSrcFieldSpecs();
+ List<SrcField> srcFields = new ArrayList<>();
+ for (SrcFieldSpec srcFieldSpec : srcFieldSpecs) {
+ if (srcFieldSpec.isBinary()) {
+ srcFields.add(new SrcField(srcFieldSpec, completeData));
+ } else {
+ log.warn(
+ "Non-binary field {} configured in binary format
data",
+ srcFieldSpec.getFieldName());
+ }
}
+ fieldValues[i++] = srcFields.isEmpty() ? null : new
MultimodalFieldValue(srcFields);
}
try {
return model.vectorization(fieldValues).toArray();
} catch (Exception e) {
- throw new RuntimeException(
- "Failed to vectorize binary data for file: " +
inputRow.toString(), e);
+ throw new RuntimeException("Failed to vectorize binary data for
file: " + inputRow, e);
}
}
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcField.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcField.java
new file mode 100644
index 0000000000..908ad5aad1
--- /dev/null
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcField.java
@@ -0,0 +1,84 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.seatunnel.transform.nlpmodel.embedding;
+
+import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
+
+import lombok.Data;
+import lombok.extern.slf4j.Slf4j;
+
+import java.io.Serializable;
+import java.util.Base64;
+
+@Data
+@Slf4j
+public class SrcField implements Serializable {
+
+ private static final long serialVersionUID = 1L;
+
+ private SrcFieldSpec fieldSpec;
+
+ private Object fieldValue;
+
+ public SrcField(SrcFieldSpec spec, Object value) {
+ // create a new object avoid to mutate original src field spec
+ this.fieldSpec =
+ new SrcFieldSpec(
+ spec.getFieldName(),
+ spec.getModalityType(),
+ spec.getPayloadFormat(),
+ spec.isModalityTypeExplicitlyConfigured());
+ this.fieldValue = value;
+ determineModalityType();
+ }
+
+ /**
+ * Determine the actual modality type based on field spec and value. The
configured modality
+ * type is always respected when it was explicitly provided by the user.
Auto-detection from the
+ * value suffix only happens for URL payloads whose modality type was not
explicitly configured,
+ * so that plain text values are never misclassified as image/video by
their content.
+ */
+ private void determineModalityType() {
+ if (fieldSpec.isModalityTypeExplicitlyConfigured() ||
!fieldSpec.isUrl()) {
+ return;
+ }
+ if (fieldValue != null) {
+ String valueStr = fieldValue.toString();
+ ModalityType detectedType = ModalityType.fromFileSuffix(valueStr);
+ if (detectedType != null) {
+ log.debug(
+ "Auto-detected modality type '{}' from value: {}",
detectedType, valueStr);
+ fieldSpec.setModalityType(detectedType);
+ }
+ }
+ }
+
+ public String toBase64() {
+ if (fieldSpec == null || !fieldSpec.isBinary()) {
+ throw new IllegalArgumentException("Payload format must be
binary");
+ }
+ if (fieldValue == null) {
+ throw new IllegalArgumentException("Binary data cannot be null or
empty");
+ }
+ if (fieldValue instanceof byte[]) {
+ return Base64.getEncoder().encodeToString((byte[]) fieldValue);
+ } else {
+ return
Base64.getEncoder().encodeToString(fieldValue.toString().getBytes());
+ }
+ }
+}
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/FieldSpec.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcFieldSpec.java
similarity index 65%
rename from
seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/FieldSpec.java
rename to
seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcFieldSpec.java
index 94ee65329e..cf5a6ae5c8 100644
---
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/FieldSpec.java
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/SrcFieldSpec.java
@@ -26,7 +26,7 @@ import java.io.Serializable;
import java.util.Map;
@Data
-public class FieldSpec implements Serializable {
+public class SrcFieldSpec implements Serializable {
private static final long serialVersionUID = 1L;
@@ -34,51 +34,31 @@ public class FieldSpec implements Serializable {
private ModalityType modalityType;
private PayloadFormat payloadFormat;
- public FieldSpec(String fieldName) {
- this.fieldName = fieldName;
- this.modalityType = ModalityType.TEXT;
- this.payloadFormat = PayloadFormat.TEXT;
- }
-
- public FieldSpec(Map.Entry<String, Object> fieldConfig) {
- String outputFieldName = fieldConfig.getKey();
- if (outputFieldName == null) {
- throw new IllegalArgumentException("Field spec cannot be null");
- }
- Object fieldValue = fieldConfig.getValue();
- try {
- if (fieldValue instanceof String) {
- parseBasicFieldSpec((String) fieldValue);
- } else {
- Map<String, Object> fieldSpecConfig = (Map<String, Object>)
fieldValue;
- parseMultimodalFieldSpec(fieldSpecConfig);
- }
- } catch (Exception e) {
- String errorMessage =
- String.format(
- "Invalid field spec for output field '%s': %s",
- outputFieldName, fieldConfig);
- throw new IllegalArgumentException(errorMessage, e);
- }
- }
+ /**
+ * Whether the modality type was explicitly configured by the user. When
false, the actual
+ * modality type can be auto-detected from the runtime value suffix; when
true, the configured
+ * modality type must be respected and never overridden.
+ */
+ private boolean modalityTypeExplicitlyConfigured;
/** Parse basic field spec: just the field name, defaults to TEXT modality
and default format */
- private void parseBasicFieldSpec(String fieldSpec) {
- if (fieldSpec == null || fieldSpec.trim().isEmpty()) {
- throw new IllegalArgumentException("Field spec cannot be null or
empty");
+ public SrcFieldSpec(String fieldName) {
+ if (fieldName == null || fieldName.trim().isEmpty()) {
+ throw new IllegalArgumentException("Field name cannot be null or
empty");
}
- this.fieldName = fieldSpec.trim();
+ this.fieldName = fieldName.trim();
this.modalityType = ModalityType.TEXT;
this.payloadFormat = PayloadFormat.TEXT;
+ this.modalityTypeExplicitlyConfigured = false;
}
/**
* Parse multimodal field spec: field name, modality, and format Supports
both formats: 1.
* Separate modality and format
*/
- private void parseMultimodalFieldSpec(Map<String, Object> fieldConfig) {
+ public SrcFieldSpec(Map<String, Object> fieldConfig) {
if (fieldConfig == null || fieldConfig.isEmpty()) {
- throw new IllegalArgumentException("Field configuration cannot be
null or empty");
+ throw new IllegalArgumentException("Field config cannot be null or
empty");
}
Object fieldNameObj = fieldConfig.get("field");
@@ -94,12 +74,14 @@ public class FieldSpec implements Serializable {
Object modalityObj = fieldConfig.get("modality");
if (modalityObj != null) {
this.modalityType = ModalityType.ofName(modalityObj.toString());
+ this.modalityTypeExplicitlyConfigured = true;
Object formatObj = fieldConfig.get("format");
if (formatObj != null) {
this.payloadFormat =
PayloadFormat.ofName(formatObj.toString());
}
} else {
this.modalityType = ModalityType.TEXT;
+ this.modalityTypeExplicitlyConfigured = false;
Object formatObj = fieldConfig.get("format");
if (formatObj != null) {
this.payloadFormat =
PayloadFormat.ofName(formatObj.toString());
@@ -109,11 +91,22 @@ public class FieldSpec implements Serializable {
}
}
- public boolean isMultimodalField() {
- return !ModalityType.TEXT.equals(modalityType);
+ public SrcFieldSpec(
+ String fieldName,
+ ModalityType modalityType,
+ PayloadFormat payloadFormat,
+ boolean modalityTypeExplicitlyConfigured) {
+ this.fieldName = fieldName;
+ this.modalityType = modalityType;
+ this.payloadFormat = payloadFormat;
+ this.modalityTypeExplicitlyConfigured =
modalityTypeExplicitlyConfigured;
}
public boolean isBinary() {
return PayloadFormat.BINARY.equals(payloadFormat);
}
+
+ public boolean isUrl() {
+ return PayloadFormat.URL.equals(payloadFormat);
+ }
}
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/VectorFieldSpec.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/VectorFieldSpec.java
new file mode 100644
index 0000000000..bf8bd937f0
--- /dev/null
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/VectorFieldSpec.java
@@ -0,0 +1,101 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.seatunnel.transform.nlpmodel.embedding;
+
+import org.apache.seatunnel.shade.org.apache.commons.lang3.StringUtils;
+
+import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
+
+import lombok.Data;
+
+import java.io.Serializable;
+import java.util.ArrayList;
+import java.util.List;
+import java.util.Map;
+import java.util.Objects;
+
+@Data
+public class VectorFieldSpec implements Serializable {
+
+ private static final long serialVersionUID = 1L;
+
+ private String fieldName;
+
+ private List<SrcFieldSpec> srcFieldSpecs;
+
+ public VectorFieldSpec(Map.Entry<String, Object> fieldConfig) {
+ this.fieldName = fieldConfig.getKey();
+ if (StringUtils.isBlank(fieldName)) {
+ throw new IllegalArgumentException("Field config name cannot be
null or empty");
+ }
+ Object fieldConfigValue = fieldConfig.getValue();
+ if (fieldConfigValue == null) {
+ throw new IllegalArgumentException(
+ "Field config value cannot be null for field: " +
fieldName);
+ }
+
+ srcFieldSpecs = new ArrayList<>();
+ try {
+ if (fieldConfigValue instanceof String) {
+ srcFieldSpecs.add(new SrcFieldSpec((String) fieldConfigValue));
+ } else if (fieldConfigValue instanceof Map) {
+ srcFieldSpecs.add(new SrcFieldSpec((Map<String, Object>)
fieldConfigValue));
+ } else {
+ List<Object> fieldConfigValues = (List<Object>)
fieldConfigValue;
+ for (Object fieldConfigValueItem : fieldConfigValues) {
+ if (fieldConfigValueItem instanceof String) {
+ srcFieldSpecs.add(new SrcFieldSpec((String)
fieldConfigValueItem));
+ } else if (fieldConfigValueItem instanceof Map) {
+ srcFieldSpecs.add(
+ new SrcFieldSpec((Map<String, Object>)
fieldConfigValueItem));
+ } else {
+ String errorMessage =
+ String.format(
+ "Invalid field spec for output field
'%s': %s",
+ fieldName, fieldConfig);
+ throw new IllegalArgumentException(errorMessage);
+ }
+ }
+ }
+ } catch (Exception e) {
+ String errorMessage =
+ String.format(
+ "Invalid field spec for output field '%s': %s",
fieldName, fieldConfig);
+ throw new IllegalArgumentException(errorMessage, e);
+ }
+ }
+
+ public boolean isMultimodalField() {
+ return srcFieldSpecs.size() > 1
+ || srcFieldSpecs.stream()
+ .anyMatch(f ->
!ModalityType.TEXT.equals(f.getModalityType()));
+ }
+
+ @Override
+ public boolean equals(Object object) {
+ if (this == object) return true;
+ if (object == null || getClass() != object.getClass()) return false;
+ VectorFieldSpec that = (VectorFieldSpec) object;
+ return Objects.equals(fieldName, that.fieldName);
+ }
+
+ @Override
+ public int hashCode() {
+ return Objects.hash(fieldName);
+ }
+}
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/multimodal/MultimodalFieldValue.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/multimodal/MultimodalFieldValue.java
index 01c3e50403..195db4a420 100644
---
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/multimodal/MultimodalFieldValue.java
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/multimodal/MultimodalFieldValue.java
@@ -17,54 +17,21 @@
package org.apache.seatunnel.transform.nlpmodel.embedding.multimodal;
-import org.apache.seatunnel.transform.nlpmodel.embedding.FieldSpec;
+import org.apache.seatunnel.transform.nlpmodel.embedding.SrcField;
import lombok.Getter;
-import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
-import java.util.Base64;
+import java.util.List;
-@Slf4j
@Getter
public class MultimodalFieldValue implements Serializable {
private static final long serialVersionUID = 1L;
- private final FieldSpec fieldSpec;
- private final Object value;
+ private final List<SrcField> srcFields;
- public MultimodalFieldValue(FieldSpec fieldSpec, Object value) {
- this.value = value;
- fieldSpec.setModalityType(determineModalityType(fieldSpec, value));
- this.fieldSpec = fieldSpec;
- }
-
- /**
- * Determine the actual modality type based on field spec and value If not
binary format,
- * analyze the value suffix to determine modality type
- */
- private ModalityType determineModalityType(FieldSpec fieldSpec, Object
value) {
-
- if (fieldSpec.isBinary()) {
- return fieldSpec.getModalityType();
- }
- if (value != null) {
- String valueStr = value.toString();
- ModalityType detectedType = ModalityType.fromFileSuffix(valueStr);
- if (detectedType != null) {
- log.debug(
- "Auto-detected modality type '{}' from value: {}",
detectedType, valueStr);
- return detectedType;
- }
- }
- return fieldSpec.getModalityType();
- }
-
- public String toBase64() {
- if (value == null) {
- throw new IllegalArgumentException("Binary data cannot be null or
empty");
- }
- return Base64.getEncoder().encodeToString(value.toString().getBytes());
+ public MultimodalFieldValue(List<SrcField> srcFields) {
+ this.srcFields = srcFields;
}
}
diff --git
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/remote/doubao/DoubaoModel.java
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/remote/doubao/DoubaoModel.java
index e02bbe74df..f175b7c9b2 100644
---
a/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/remote/doubao/DoubaoModel.java
+++
b/seatunnel-transforms-v2/src/main/java/org/apache/seatunnel/transform/nlpmodel/embedding/remote/doubao/DoubaoModel.java
@@ -28,7 +28,8 @@ import
org.apache.seatunnel.transform.nlpmodel.ModelInvocationErrorType;
import org.apache.seatunnel.transform.nlpmodel.ModelInvocationException;
import org.apache.seatunnel.transform.nlpmodel.ModelInvocationOptions;
import org.apache.seatunnel.transform.nlpmodel.ProviderAdapter;
-import org.apache.seatunnel.transform.nlpmodel.embedding.FieldSpec;
+import org.apache.seatunnel.transform.nlpmodel.embedding.SrcField;
+import org.apache.seatunnel.transform.nlpmodel.embedding.SrcFieldSpec;
import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.MultimodalFieldValue;
import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.MultimodalModel;
@@ -45,6 +46,7 @@ import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.Arrays;
+import java.util.Collections;
import java.util.List;
public class DoubaoModel extends MultimodalModel {
@@ -143,7 +145,7 @@ public class DoubaoModel extends MultimodalModel {
public List<List<Float>> multimodalVector(Object[] fields) throws
IOException {
if (singleVectorizedInputNumber > 1) {
throw new IllegalArgumentException(
- "Doubao does not support batch multimodal vectorization in
a single request. ");
+ "Doubao does not support batch multimodal vectorization in
a single request.");
}
List<List<Float>> vectors = new ArrayList<>();
for (Object field : fields) {
@@ -188,7 +190,10 @@ public class DoubaoModel extends MultimodalModel {
? multimodalVector(
new Object[] {
new MultimodalFieldValue(
- new FieldSpec(DIMENSION_EXAMPLE),
DIMENSION_EXAMPLE)
+ Collections.singletonList(
+ new SrcField(
+ new
SrcFieldSpec(DIMENSION_EXAMPLE),
+
DIMENSION_EXAMPLE)))
})
.get(0)
.size()
@@ -355,35 +360,40 @@ public class DoubaoModel extends MultimodalModel {
ObjectNode requestNode = OBJECT_MAPPER.createObjectNode();
requestNode.put("model", model);
requestNode.put("encoding_format", "float");
- ArrayNode inputDatas = OBJECT_MAPPER.createArrayNode();
- inputDatas.add(inputRawData(field));
- requestNode.set("input", inputDatas);
+ ArrayNode inputNode = OBJECT_MAPPER.createArrayNode();
+ inputNode.addAll(inputRawData(field));
+ requestNode.set("input", inputNode);
return requestNode;
}
- protected ObjectNode inputRawData(MultimodalFieldValue field) {
- ObjectNode rawDataNode = OBJECT_MAPPER.createObjectNode();
- FieldSpec fieldSpec = field.getFieldSpec();
- String fieldValue = field.getValue().toString().trim();
- ModalityType fieldSpecModalityType = fieldSpec.getModalityType();
- String modalityParamName = getModalityParamName(fieldSpecModalityType);
- rawDataNode.put("type", modalityParamName);
- if (ModalityType.TEXT == fieldSpecModalityType) {
- rawDataNode.put(modalityParamName, fieldValue);
- return rawDataNode;
- }
+ protected List<ObjectNode> inputRawData(MultimodalFieldValue field) {
+ List<ObjectNode> rawDataNodes = new ArrayList<>();
+ List<SrcField> srcFields = field.getSrcFields();
+ for (SrcField srcField : srcFields) {
+ ObjectNode rawDataNode = OBJECT_MAPPER.createObjectNode();
+ String fieldValue = srcField.getFieldValue().toString().trim();
+ ModalityType fieldSpecModalityType =
srcField.getFieldSpec().getModalityType();
+ String modalityParamName =
getModalityParamName(fieldSpecModalityType);
+ rawDataNode.put("type", modalityParamName);
+ if (ModalityType.TEXT == fieldSpecModalityType) {
+ rawDataNode.put(modalityParamName, fieldValue);
+ rawDataNodes.add(rawDataNode);
+ continue;
+ }
- if (fieldSpec.isBinary()) {
- fieldValue =
- String.format(
- BASE64_PARAM_TEMPLATE,
-
fieldSpecModalityType.getGroup().name().toLowerCase(),
- fieldSpecModalityType.getName(),
- field.toBase64());
+ if (srcField.getFieldSpec().isBinary()) {
+ fieldValue =
+ String.format(
+ BASE64_PARAM_TEMPLATE,
+
fieldSpecModalityType.getGroup().name().toLowerCase(),
+ fieldSpecModalityType.getName(),
+ srcField.toBase64());
+ }
+ rawDataNode.set(
+ modalityParamName,
OBJECT_MAPPER.createObjectNode().put("url", fieldValue));
+ rawDataNodes.add(rawDataNode);
}
- rawDataNode.set(modalityParamName,
OBJECT_MAPPER.createObjectNode().put("url", fieldValue));
-
- return rawDataNode;
+ return rawDataNodes;
}
private String getModalityParamName(ModalityType inputType) {
diff --git
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/DoubaoMultimodalModelTest.java
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/DoubaoMultimodalModelTest.java
index b9ae009e8a..0489a76ace 100644
---
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/DoubaoMultimodalModelTest.java
+++
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/DoubaoMultimodalModelTest.java
@@ -17,40 +17,59 @@
package org.apache.seatunnel.transform.embedding;
-import org.apache.seatunnel.shade.com.fasterxml.jackson.databind.ObjectMapper;
import
org.apache.seatunnel.shade.com.fasterxml.jackson.databind.node.ObjectNode;
-import org.apache.seatunnel.transform.nlpmodel.embedding.FieldSpec;
+import org.apache.seatunnel.transform.nlpmodel.embedding.SrcField;
+import org.apache.seatunnel.transform.nlpmodel.embedding.VectorFieldSpec;
import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.MultimodalFieldValue;
import
org.apache.seatunnel.transform.nlpmodel.embedding.remote.doubao.DoubaoModel;
+import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Assertions;
+import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.io.IOException;
+import java.util.Arrays;
+import java.util.Base64;
+import java.util.Collections;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
+import java.util.concurrent.ThreadLocalRandom;
public class DoubaoMultimodalModelTest {
- private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper();
+ private DoubaoModel model;
- @Test
- void testMultimodalBodyWithText() throws IOException {
- DoubaoModel model =
+ @BeforeEach
+ void setUp() {
+ this.model =
new DoubaoModel(
"test-api-key",
"doubao-embedding-vision",
"https://ark.cn-beijing.volces.com/api/v3/embeddings",
1);
+ }
+ @AfterEach
+ void tearDown() throws IOException {
+ if (model != null) {
+ model.close();
+ }
+ }
+
+ @Test
+ void testMultimodalBodyWithText() {
Map.Entry<String, Object> textFieldEntry =
- new java.util.AbstractMap.SimpleEntry<>("text_vector", "Hello
world");
- FieldSpec fieldSpec = new FieldSpec(textFieldEntry);
+ new java.util.AbstractMap.SimpleEntry<>("text_vector",
"text_field");
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(textFieldEntry);
MultimodalFieldValue multimodalFieldValue =
- new MultimodalFieldValue(fieldSpec, "Hello world");
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0), "Hello world")));
ObjectNode result = model.multimodalBody(multimodalFieldValue);
@@ -63,36 +82,29 @@ public class DoubaoMultimodalModelTest {
Assertions.assertEquals("Hello world", inputNode.get("text").asText());
Assertions.assertFalse(inputNode.has("image_url"));
Assertions.assertFalse(inputNode.has("video_url"));
-
- model.close();
}
/**
- * { "model" : "doubao-embedding-vision", "encoding_format" : "float",
"input" : [ { "type" :
- * "image_url", "image_url" : { "url" :
- *
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg" } }]
}
+ * { "model": "doubao-embedding-vision", "encoding_format": "float",
"input": [ { "type":
+ * "image_url", "image_url": { "url":
+ *
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg" } }
] }
*/
@Test
- void testMultimodalBodyWithImage() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
-
+ void testMultimodalBodyWithImage() {
Map<String, Object> imageFieldConfig = new HashMap<>();
imageFieldConfig.put("field", "image_field");
imageFieldConfig.put("modality", "jpeg");
imageFieldConfig.put("format", "url");
-
Map.Entry<String, Object> imageFieldEntry =
new java.util.AbstractMap.SimpleEntry<>("image_vector",
imageFieldConfig);
- FieldSpec fieldSpec = new FieldSpec(imageFieldEntry);
+
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(imageFieldEntry);
MultimodalFieldValue multimodalFieldValue =
new MultimodalFieldValue(
- fieldSpec,
-
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg");
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0),
+
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg")));
ObjectNode result = model.multimodalBody(multimodalFieldValue);
@@ -110,33 +122,28 @@ public class DoubaoMultimodalModelTest {
inputNode.get("image_url").get("url").asText());
Assertions.assertFalse(inputNode.has("text"));
Assertions.assertFalse(inputNode.has("video_url"));
-
- model.close();
}
/**
- * { "model" : "doubao-embedding-vision", "encoding_format" : "float",
"input" : [ { "type" :
- * "video_url", "video_url" : { "url" : "https://example.com/video.mp4" }
} ] }
+ * { "model": "doubao-embedding-vision", "encoding_format": "float",
"input": [ { "type":
+ * "video_url", "video_url": { "url": "https://example.com/video.mp4" } }
] }
*/
@Test
- void testMultimodalBodyWithVideo() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
-
+ void testMultimodalBodyWithVideo() {
Map<String, Object> videoFieldConfig = new HashMap<>();
videoFieldConfig.put("field", "video_field");
videoFieldConfig.put("modality", "mP4");
videoFieldConfig.put("format", "url");
-
Map.Entry<String, Object> videoFieldEntry =
new java.util.AbstractMap.SimpleEntry<>("video_vector",
videoFieldConfig);
- FieldSpec fieldSpec = new FieldSpec(videoFieldEntry);
+
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(videoFieldEntry);
MultimodalFieldValue multimodalFieldValue =
- new MultimodalFieldValue(fieldSpec,
"https://example.com/video.mp4");
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0),
+ "https://example.com/video.mp4")));
ObjectNode result = model.multimodalBody(multimodalFieldValue);
@@ -151,8 +158,6 @@ public class DoubaoMultimodalModelTest {
"https://example.com/video.mp4",
inputNode.get("video_url").get("url").asText());
Assertions.assertFalse(inputNode.has("text"));
Assertions.assertFalse(inputNode.has("image_url"));
-
- model.close();
}
/**
@@ -160,50 +165,149 @@ public class DoubaoMultimodalModelTest {
* f"data:image/<IMAGE_FORMAT>;base64,{base64_image}" } }
*/
@Test
- void testMultimodalBodyWithBinaryImage() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision-250615",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
-
+ void testMultimodalBodyWithBinaryImage() {
Map<String, Object> binaryImageFieldConfig = new HashMap<>();
binaryImageFieldConfig.put("field", "binary_image_field");
binaryImageFieldConfig.put("modality", "png");
binaryImageFieldConfig.put("format", "binary");
-
Map.Entry<String, Object> binaryImageFieldEntry =
new java.util.AbstractMap.SimpleEntry<>(
"binary_image_vector", binaryImageFieldConfig);
- FieldSpec fieldSpec = new FieldSpec(binaryImageFieldEntry);
+ VectorFieldSpec vectorFieldSpec = new
VectorFieldSpec(binaryImageFieldEntry);
byte[] mockImageData =
"mock-image-data".getBytes(java.nio.charset.StandardCharsets.UTF_8);
MultimodalFieldValue multimodalFieldValue =
- new MultimodalFieldValue(fieldSpec, mockImageData);
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0), mockImageData)));
ObjectNode result = model.multimodalBody(multimodalFieldValue);
-
- Assertions.assertEquals("doubao-embedding-vision-250615",
result.get("model").asText());
+ Assertions.assertEquals("doubao-embedding-vision",
result.get("model").asText());
Assertions.assertEquals("float",
result.get("encoding_format").asText());
Assertions.assertEquals(1, result.get("input").size());
ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
Assertions.assertEquals("image_url", inputNode.get("type").asText());
Assertions.assertTrue(inputNode.has("image_url"));
+ Assertions.assertTrue(
+ inputNode
+ .get("image_url")
+ .get("url")
+ .asText()
+
.endsWith(Base64.getEncoder().encodeToString(mockImageData)));
+ }
+
+ /**
+ * { "model": "doubao-embedding-vision", "encoding_format": "float",
"input": [ { "type":
+ * "text", "text": "Hello world 1" }, { "type": "text", "text": "Hello
world 2" } ] }
+ */
+ @Test
+ void testMultimodalBodyWithSameModalityList() {
+ Map.Entry<String, Object> vectorFieldEntry =
+ new java.util.AbstractMap.SimpleEntry<>(
+ "same_multimodal_vector",
Arrays.asList("text_field_1", "text_field_2"));
+ VectorFieldSpec vectorFieldSpec = new
VectorFieldSpec(vectorFieldEntry);
+ MultimodalFieldValue multimodalFieldValue =
+ new MultimodalFieldValue(
+ Arrays.asList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0), "Hello world 1"),
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(1),
+ "Hello world 2")));
+
+ ObjectNode result = model.multimodalBody(multimodalFieldValue);
+ Assertions.assertEquals("doubao-embedding-vision",
result.get("model").asText());
+ Assertions.assertEquals("float",
result.get("encoding_format").asText());
+ Assertions.assertEquals(2, result.get("input").size());
- model.close();
+ ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
+ Assertions.assertEquals("text", inputNode.get("type").asText());
+ Assertions.assertEquals("Hello world 1",
inputNode.get("text").asText());
+ Assertions.assertFalse(inputNode.has("image_url"));
+ Assertions.assertFalse(inputNode.has("video_url"));
+
+ inputNode = (ObjectNode) result.get("input").get(1);
+ Assertions.assertEquals("text", inputNode.get("type").asText());
+ Assertions.assertEquals("Hello world 2",
inputNode.get("text").asText());
+ Assertions.assertFalse(inputNode.has("image_url"));
+ Assertions.assertFalse(inputNode.has("video_url"));
}
+ /**
+ * { "model": "doubao-embedding-vision", "encoding_format": "float",
"input": [ { "type":
+ * "text", "text": "Hello world" }, { "type": "image_url", "image_url": {
"url":
+ *
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg" } },
{ "type":
+ * "video_url", "video_url": { "url": "https://example.com/video.mp4" } }
] }
+ */
@Test
- void testParseMultimodalVectorResponseSuccess() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
+ void testMultimodalBodyWithDifferentModalityList() {
+ Object textFieldConfig = "text_field";
+ if (ThreadLocalRandom.current().nextBoolean()) {
+ Map<String, Object> textFieldConfigMap = new HashMap<>();
+ textFieldConfigMap.put("field", "text_field");
+ textFieldConfigMap.put("modality", "text");
+ textFieldConfigMap.put("format", "text");
+ textFieldConfig = textFieldConfigMap;
+ }
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "jpeg");
+ imageFieldConfig.put("format", "url");
+ Map<String, Object> videoFieldConfig = new HashMap<>();
+ videoFieldConfig.put("field", "video_field");
+ videoFieldConfig.put("modality", "mp4");
+ videoFieldConfig.put("format", "url");
+ Map.Entry<String, Object> vectorFieldEntry =
+ new java.util.AbstractMap.SimpleEntry<>(
+ "different_multimodal_vector",
+ Arrays.asList(textFieldConfig, imageFieldConfig,
videoFieldConfig));
+
+ VectorFieldSpec vectorFieldSpec = new
VectorFieldSpec(vectorFieldEntry);
+ MultimodalFieldValue multimodalFieldValue =
+ new MultimodalFieldValue(
+ Arrays.asList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0), "Hello world"),
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(1),
+
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg"),
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(2),
+ "https://example.com/video.mp4")));
+ ObjectNode result = model.multimodalBody(multimodalFieldValue);
+ Assertions.assertEquals("doubao-embedding-vision",
result.get("model").asText());
+ Assertions.assertEquals("float",
result.get("encoding_format").asText());
+ Assertions.assertEquals(3, result.get("input").size());
+
+ ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
+ Assertions.assertEquals("text", inputNode.get("type").asText());
+ Assertions.assertEquals("Hello world", inputNode.get("text").asText());
+ Assertions.assertFalse(inputNode.has("image_url"));
+ Assertions.assertFalse(inputNode.has("video_url"));
+
+ inputNode = (ObjectNode) result.get("input").get(1);
+ Assertions.assertEquals("image_url", inputNode.get("type").asText());
+ Assertions.assertTrue(inputNode.has("image_url"));
+ Assertions.assertEquals(
+
"https://ck-test.tos-cn-beijing.volces.com/vlm/pexels-photo-27163466.jpeg",
+ inputNode.get("image_url").get("url").asText());
+ Assertions.assertFalse(inputNode.has("text"));
+ Assertions.assertFalse(inputNode.has("video_url"));
+
+ inputNode = (ObjectNode) result.get("input").get(2);
+ Assertions.assertEquals("video_url", inputNode.get("type").asText());
+ Assertions.assertTrue(inputNode.has("video_url"));
+ Assertions.assertEquals(
+ "https://example.com/video.mp4",
inputNode.get("video_url").get("url").asText());
+ Assertions.assertFalse(inputNode.has("text"));
+ Assertions.assertFalse(inputNode.has("image_url"));
+ }
+
+ @Test
+ void testParseMultimodalVectorResponseSuccess() throws IOException {
String successResponse =
"{\n"
+ " \"created\": 1743575029,\n"
@@ -236,79 +340,136 @@ public class DoubaoMultimodalModelTest {
Assertions.assertEquals(-0.318359375f, result.get(2), 0.0001f);
Assertions.assertEquals(0.255859375f, result.get(3), 0.0001f);
Assertions.assertEquals(1.5f, result.get(4), 0.0001f);
-
- model.close();
}
@Test
- void testUrlAutoDetectModality() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
-
+ void testUrlAutoDetectModality() {
+ // Explicitly configured modality (png) must be respected and NOT
overridden by the runtime
+ // value suffix (.jpg).
Map<String, Object> fieldConfig = new HashMap<>();
fieldConfig.put("field", "image_field");
fieldConfig.put("format", "url");
fieldConfig.put("modality", "png");
- Map.Entry<String, Object> fieldEntry =
+ Map.Entry<String, Object> imageFieldEntry =
new java.util.AbstractMap.SimpleEntry<>("image_vector",
fieldConfig);
- FieldSpec fieldSpec = new FieldSpec(fieldEntry);
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(imageFieldEntry);
MultimodalFieldValue multimodalFieldValue =
- new MultimodalFieldValue(fieldSpec,
"https://example.com/photo.jpg");
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0),
+ "https://example.com/photo.jpg")));
Assertions.assertEquals(
- ModalityType.JPEG,
multimodalFieldValue.getFieldSpec().getModalityType());
+ ModalityType.PNG,
+
multimodalFieldValue.getSrcFields().get(0).getFieldSpec().getModalityType());
ObjectNode result = model.multimodalBody(multimodalFieldValue);
ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
Assertions.assertEquals("image_url", inputNode.get("type").asText());
+ // No modality configured -> auto-detect from the value suffix (.jpg
-> jpeg).
Map<String, Object> fieldConfig2 = new HashMap<>();
fieldConfig2.put("field", "image_field");
fieldConfig2.put("format", "url");
- fieldEntry = new java.util.AbstractMap.SimpleEntry<>("image_vector",
fieldConfig2);
- fieldSpec = new FieldSpec(fieldEntry);
-
- multimodalFieldValue = new MultimodalFieldValue(fieldSpec,
"https://example.com/photo.jpg");
+ imageFieldEntry = new
java.util.AbstractMap.SimpleEntry<>("image_vector", fieldConfig2);
+ vectorFieldSpec = new VectorFieldSpec(imageFieldEntry);
+ multimodalFieldValue =
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0),
+ "https://example.com/photo.jpg")));
Assertions.assertEquals(
- ModalityType.JPEG,
multimodalFieldValue.getFieldSpec().getModalityType());
+ ModalityType.JPEG,
+
multimodalFieldValue.getSrcFields().get(0).getFieldSpec().getModalityType());
result = model.multimodalBody(multimodalFieldValue);
inputNode = (ObjectNode) result.get("input").get(0);
Assertions.assertEquals("image_url", inputNode.get("type").asText());
+ }
- model.close();
+ @Test
+ void testExplicitModalityNotOverriddenBySuffix() {
+ // Regression: modality = png + runtime value photo.jpg should stay
png.
+ Map<String, Object> fieldConfig = new HashMap<>();
+ fieldConfig.put("field", "image_field");
+ fieldConfig.put("format", "url");
+ fieldConfig.put("modality", "png");
+ Map.Entry<String, Object> imageFieldEntry =
+ new java.util.AbstractMap.SimpleEntry<>("image_vector",
fieldConfig);
+
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(imageFieldEntry);
+ SrcField srcField =
+ new SrcField(
+ vectorFieldSpec.getSrcFieldSpecs().get(0),
"https://example.com/photo.jpg");
+
+ Assertions.assertEquals(ModalityType.PNG,
srcField.getFieldSpec().getModalityType());
+
Assertions.assertTrue(srcField.getFieldSpec().isModalityTypeExplicitlyConfigured());
}
@Test
- void testBinaryAutoDetectModality() throws IOException {
- DoubaoModel model =
- new DoubaoModel(
- "test-api-key",
- "doubao-embedding-vision",
- "https://ark.cn-beijing.volces.com/api/v3/embeddings",
- 1);
+ void testMixedConfigTextFieldWithImageSuffixStaysText() {
+ // Regression: in a mixed multimodal job, a plain text field whose
value happens to end with
+ // a known image suffix (foo.jpg) must NOT be misclassified as an
image.
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "jpeg");
+ imageFieldConfig.put("format", "url");
+
+ Map.Entry<String, Object> entry =
+ new java.util.AbstractMap.SimpleEntry<>(
+ "mix_vector", Arrays.asList("text_field",
imageFieldConfig));
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+
+ // first src field is the plain text field, with a value that ends
with .jpg
+ SrcField textSrcField =
+ new SrcField(vectorFieldSpec.getSrcFieldSpecs().get(0), "this
is foo.jpg");
+ Assertions.assertEquals(ModalityType.TEXT,
textSrcField.getFieldSpec().getModalityType());
+
Assertions.assertFalse(textSrcField.getFieldSpec().isModalityTypeExplicitlyConfigured());
+
+ MultimodalFieldValue multimodalFieldValue =
+ new
MultimodalFieldValue(Collections.singletonList(textSrcField));
+ ObjectNode result = model.multimodalBody(multimodalFieldValue);
+ ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
+ Assertions.assertEquals("text", inputNode.get("type").asText());
+ Assertions.assertEquals("this is foo.jpg",
inputNode.get("text").asText());
+ }
+
+ @Test
+ void testNoModalityPlainTextValueStaysText() {
+ // No modality configured and value has no recognizable suffix ->
stays TEXT.
+ Map.Entry<String, Object> entry =
+ new java.util.AbstractMap.SimpleEntry<>("text_vector", "hello
world");
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ SrcField srcField = new
SrcField(vectorFieldSpec.getSrcFieldSpecs().get(0), "hello world");
+
+ Assertions.assertEquals(ModalityType.TEXT,
srcField.getFieldSpec().getModalityType());
+
Assertions.assertFalse(srcField.getFieldSpec().isModalityTypeExplicitlyConfigured());
+ }
+ @Test
+ void testBinaryAutoDetectModality() {
Map<String, Object> fieldConfig = new HashMap<>();
fieldConfig.put("field", "image_field");
fieldConfig.put("format", "binary");
fieldConfig.put("modality", "png");
- Map.Entry<String, Object> fieldEntry =
+ Map.Entry<String, Object> imageFieldEntry =
new java.util.AbstractMap.SimpleEntry<>("image_vector",
fieldConfig);
- FieldSpec fieldSpec = new FieldSpec(fieldEntry);
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(imageFieldEntry);
MultimodalFieldValue multimodalFieldValue =
- new MultimodalFieldValue(fieldSpec,
"https://example.com/photo.jpg");
+ new MultimodalFieldValue(
+ Collections.singletonList(
+ new SrcField(
+
vectorFieldSpec.getSrcFieldSpecs().get(0),
+ "https://example.com/photo.jpg")));
Assertions.assertEquals(
- ModalityType.PNG,
multimodalFieldValue.getFieldSpec().getModalityType());
+ ModalityType.PNG,
+
multimodalFieldValue.getSrcFields().get(0).getFieldSpec().getModalityType());
ObjectNode result = model.multimodalBody(multimodalFieldValue);
ObjectNode inputNode = (ObjectNode) result.get("input").get(0);
Assertions.assertEquals("image_url", inputNode.get("type").asText());
-
- model.close();
}
}
diff --git
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/FieldSpecTest.java
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/FieldSpecTest.java
deleted file mode 100644
index c97372f8fe..0000000000
---
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/FieldSpecTest.java
+++ /dev/null
@@ -1,114 +0,0 @@
-/*
- * Licensed to the Apache Software Foundation (ASF) under one or more
- * contributor license agreements. See the NOTICE file distributed with
- * this work for additional information regarding copyright ownership.
- * The ASF licenses this file to You under the Apache License, Version 2.0
- * (the "License"); you may not use this file except in compliance with
- * the License. You may obtain a copy of the License at
- *
- * http://www.apache.org/licenses/LICENSE-2.0
- *
- * Unless required by applicable law or agreed to in writing, software
- * distributed under the License is distributed on an "AS IS" BASIS,
- * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
- * See the License for the specific language governing permissions and
- * limitations under the License.
- */
-
-package org.apache.seatunnel.transform.embedding;
-
-import org.apache.seatunnel.transform.nlpmodel.embedding.FieldSpec;
-import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
-import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.PayloadFormat;
-
-import org.junit.jupiter.api.Assertions;
-import org.junit.jupiter.api.Test;
-
-import java.util.AbstractMap;
-import java.util.HashMap;
-import java.util.Map;
-
-public class FieldSpecTest {
-
- @Test
- void testMapEntryConstructorWithStringValue() {
- Map.Entry<String, Object> entry =
- new AbstractMap.SimpleEntry<>("book_intro_vector",
"book_intro");
- FieldSpec fieldSpec = new FieldSpec(entry);
- Assertions.assertEquals("book_intro", fieldSpec.getFieldName());
- Assertions.assertEquals(ModalityType.TEXT,
fieldSpec.getModalityType());
- Assertions.assertEquals(PayloadFormat.TEXT,
fieldSpec.getPayloadFormat());
- Assertions.assertFalse(fieldSpec.isMultimodalField());
- Assertions.assertFalse(fieldSpec.isBinary());
- }
-
- @Test
- void testMapEntryConstructorWithStringValueTrimming() {
- Map.Entry<String, Object> entry =
- new AbstractMap.SimpleEntry<>("book_intro_vector", "
book_intro ");
- FieldSpec fieldSpec = new FieldSpec(entry);
- Assertions.assertEquals("book_intro", fieldSpec.getFieldName());
- Assertions.assertEquals(ModalityType.TEXT,
fieldSpec.getModalityType());
- Assertions.assertEquals(PayloadFormat.TEXT,
fieldSpec.getPayloadFormat());
- }
-
- @Test
- void testMapEntryConstructorWithNullKey() {
- Map.Entry<String, Object> entry = new AbstractMap.SimpleEntry<>(null,
"book_intro");
- IllegalArgumentException exception =
- Assertions.assertThrows(IllegalArgumentException.class, () ->
new FieldSpec(entry));
- Assertions.assertTrue(exception.getMessage().contains("Field spec
cannot be null"));
- }
-
- @Test
- void testMapEntryConstructorWithEmpty() {
- Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_intro_vector", null);
- IllegalArgumentException exception =
- Assertions.assertThrows(IllegalArgumentException.class, () ->
new FieldSpec(entry));
- Assertions.assertTrue(
- exception.getMessage().contains("Invalid field spec for output
field"));
-
- Map.Entry<String, Object> entry2 = new
AbstractMap.SimpleEntry<>("book_intro_vector", "");
- exception =
- Assertions.assertThrows(
- IllegalArgumentException.class, () -> new
FieldSpec(entry2));
- Assertions.assertTrue(
- exception.getMessage().contains("Invalid field spec for output
field"));
- }
-
- @Test
- void testMapEntryConstructorWithMapValue() {
-
- Map<String, Object> fieldConfig = new HashMap<>();
- fieldConfig.put("field", "book_image");
- fieldConfig.put("modality", "jpeg");
- fieldConfig.put("format", "binary");
-
- Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_field", fieldConfig);
-
- FieldSpec fieldSpec = new FieldSpec(entry);
-
- Assertions.assertEquals("book_image", fieldSpec.getFieldName());
- Assertions.assertEquals(ModalityType.JPEG,
fieldSpec.getModalityType());
- Assertions.assertEquals(PayloadFormat.BINARY,
fieldSpec.getPayloadFormat());
- Assertions.assertTrue(fieldSpec.isMultimodalField());
- Assertions.assertTrue(fieldSpec.isBinary());
- }
-
- @Test
- void testMapEntryConstructorWithMapValueNoModality() {
- Map<String, Object> fieldConfig = new HashMap<>();
- fieldConfig.put("field", "book_intro");
- fieldConfig.put("modality", "text");
- fieldConfig.put("format", "text");
-
- Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_field", fieldConfig);
-
- FieldSpec fieldSpec = new FieldSpec(entry);
-
- Assertions.assertEquals("book_intro", fieldSpec.getFieldName());
- Assertions.assertEquals(ModalityType.TEXT,
fieldSpec.getModalityType());
- Assertions.assertEquals(PayloadFormat.TEXT,
fieldSpec.getPayloadFormat());
- Assertions.assertFalse(fieldSpec.isMultimodalField());
- }
-}
diff --git
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/MultimodalConfigTest.java
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/MultimodalConfigTest.java
index ba5eae1f71..aecba89489 100644
---
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/MultimodalConfigTest.java
+++
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/MultimodalConfigTest.java
@@ -35,7 +35,10 @@ import org.junit.jupiter.api.Test;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashMap;
+import java.util.LinkedHashMap;
+import java.util.List;
import java.util.Map;
+import java.util.concurrent.ThreadLocalRandom;
public class MultimodalConfigTest {
@@ -44,7 +47,10 @@ public class MultimodalConfigTest {
PhysicalColumn.of("text_field", BasicType.STRING_TYPE, 255L, true,
null, ""),
PhysicalColumn.of("image_field", BasicType.STRING_TYPE, 255L,
true, null, ""),
PhysicalColumn.of("video_field", BasicType.STRING_TYPE, 255L,
true, null, ""),
- PhysicalColumn.of("mixed_field", BasicType.STRING_TYPE, 255L,
true, null, "")
+ PhysicalColumn.of("mixed_field", BasicType.STRING_TYPE, 255L,
true, null, ""),
+ PhysicalColumn.of("text_field_2", BasicType.STRING_TYPE, 255L,
true, null, ""),
+ PhysicalColumn.of("image_field_2", BasicType.STRING_TYPE, 255L,
true, null, ""),
+ PhysicalColumn.of("video_field_2", BasicType.STRING_TYPE, 255L,
true, null, ""),
};
TableSchema tableSchema =
TableSchema.builder().columns(Arrays.asList(columns)).build();
@@ -187,6 +193,110 @@ public class MultimodalConfigTest {
Assertions.assertTrue(transform.isMultimodalFields());
}
+ @Test
+ void testIsMultimodalFieldsDetectionWithMixedListFields() {
+ CatalogTable catalogTable = createTestCatalogTable();
+
+ Map<String, Object> configMap = new HashMap<>();
+ configMap.put(ModelTransformConfig.MODEL_PROVIDER.key(),
ModelProvider.DOUBAO.name());
+ configMap.put(ModelTransformConfig.MODEL.key(),
"doubao-embedding-vision");
+ configMap.put(ModelTransformConfig.API_KEY.key(), "test-api-key");
+ configMap.put(ModelTransformConfig.API_PATH.key(),
"https://api.test.com/embeddings");
+
+ Map<String, Object> vectorizationFields = new HashMap<>();
+ // Text type
+ List<Object> textFieldConfigList = Arrays.asList("text_field",
"text_field_2");
+ if (ThreadLocalRandom.current().nextBoolean()) {
+ textFieldConfigList = new ArrayList<>();
+ Map<String, Object> textFieldConfig = new HashMap<>();
+ textFieldConfig.put("field", "text_field");
+ textFieldConfig.put("modality", "text");
+ textFieldConfigList.add(textFieldConfig);
+ Map<String, Object> textFieldConfig2 = new HashMap<>();
+ textFieldConfig2.put("field", "text_field_2");
+ textFieldConfig2.put("modality", "text");
+ textFieldConfigList.add(textFieldConfig2);
+ }
+ vectorizationFields.put("text_vector", textFieldConfigList);
+
+ // Image type
+ List<Map<String, Object>> imageFieldConfigList = new ArrayList<>();
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "png");
+ imageFieldConfigList.add(imageFieldConfig);
+ Map<String, Object> imageFieldConfig2 = new HashMap<>();
+ imageFieldConfig2.put("field", "image_field_2");
+ imageFieldConfig2.put("modality", "png");
+ imageFieldConfigList.add(imageFieldConfig2);
+ vectorizationFields.put("image_vector", imageFieldConfigList);
+
+ // Video type
+ List<Map<String, Object>> videoFieldConfigList = new ArrayList<>();
+ Map<String, Object> videoFieldConfig = new HashMap<>();
+ videoFieldConfig.put("field", "video_field");
+ videoFieldConfig.put("modality", "mp4");
+ videoFieldConfig.put("format", "url");
+ videoFieldConfigList.add(videoFieldConfig);
+ Map<String, Object> videoFieldConfig2 = new HashMap<>();
+ videoFieldConfig2.put("field", "video_field_2");
+ videoFieldConfig2.put("modality", "mp4");
+ videoFieldConfig2.put("format", "url");
+ videoFieldConfigList.add(videoFieldConfig2);
+ vectorizationFields.put("video_vector", videoFieldConfigList);
+
+ configMap.put(EmbeddingTransformConfig.VECTORIZATION_FIELDS.key(),
vectorizationFields);
+
+ ReadonlyConfig config = ReadonlyConfig.fromMap(configMap);
+
+ EmbeddingTransform transform = new EmbeddingTransform(config,
catalogTable);
+ Assertions.assertNotNull(transform);
+ Assertions.assertTrue(transform.isMultimodalFields());
+ }
+
+ @Test
+ void testIsMultimodalFieldsDetectionWithMixedTypeFields() {
+ CatalogTable catalogTable = createTestCatalogTable();
+
+ Map<String, Object> configMap = new HashMap<>();
+ configMap.put(ModelTransformConfig.MODEL_PROVIDER.key(),
ModelProvider.DOUBAO.name());
+ configMap.put(ModelTransformConfig.MODEL.key(),
"doubao-embedding-vision");
+ configMap.put(ModelTransformConfig.API_KEY.key(), "test-api-key");
+ configMap.put(ModelTransformConfig.API_PATH.key(),
"https://api.test.com/embeddings");
+
+ Map<String, Object> vectorizationFields = new LinkedHashMap<>();
+ // Video type
+ Map<String, Object> videoFieldConfig = new HashMap<>();
+ videoFieldConfig.put("field", "video_field");
+ videoFieldConfig.put("modality", "mp4");
+ videoFieldConfig.put("format", "url");
+ vectorizationFields.put("video_vector", videoFieldConfig);
+
+ // Image type
+ List<Map<String, Object>> imageFieldConfigList = new ArrayList<>();
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "png");
+ imageFieldConfigList.add(imageFieldConfig);
+ Map<String, Object> imageFieldConfig2 = new HashMap<>();
+ imageFieldConfig2.put("field", "image_field_2");
+ imageFieldConfig2.put("modality", "png");
+ imageFieldConfigList.add(imageFieldConfig2);
+ vectorizationFields.put("image_vector", imageFieldConfigList);
+
+ // Text type
+ Object textFieldConfig = "text_field";
+ vectorizationFields.put("text_vector", textFieldConfig);
+
+ configMap.put(EmbeddingTransformConfig.VECTORIZATION_FIELDS.key(),
vectorizationFields);
+
+ ReadonlyConfig config = ReadonlyConfig.fromMap(configMap);
+
+ EmbeddingTransform transform = new EmbeddingTransform(config,
catalogTable);
+ Assertions.assertNotNull(transform);
+ Assertions.assertTrue(transform.isMultimodalFields());
+ }
+
@Test
void testMultimodalModelValidationFailure() {
CatalogTable catalogTable = createTestCatalogTable();
diff --git
a/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/VectorFieldSpecTest.java
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/VectorFieldSpecTest.java
new file mode 100644
index 0000000000..42677ead46
--- /dev/null
+++
b/seatunnel-transforms-v2/src/test/java/org/apache/seatunnel/transform/embedding/VectorFieldSpecTest.java
@@ -0,0 +1,200 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.seatunnel.transform.embedding;
+
+import org.apache.seatunnel.transform.nlpmodel.embedding.SrcFieldSpec;
+import org.apache.seatunnel.transform.nlpmodel.embedding.VectorFieldSpec;
+import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.ModalityType;
+import
org.apache.seatunnel.transform.nlpmodel.embedding.multimodal.PayloadFormat;
+
+import org.junit.jupiter.api.Assertions;
+import org.junit.jupiter.api.Test;
+
+import java.util.AbstractMap;
+import java.util.Arrays;
+import java.util.HashMap;
+import java.util.List;
+import java.util.Map;
+
+public class VectorFieldSpecTest {
+
+ @Test
+ void testMapEntryConstructorWithStringValue() {
+ Map.Entry<String, Object> entry =
+ new AbstractMap.SimpleEntry<>("book_intro_vector",
"book_intro");
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ Assertions.assertEquals("book_intro_vector",
vectorFieldSpec.getFieldName());
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+ Assertions.assertEquals("book_intro", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(vectorFieldSpec.isMultimodalField());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+ }
+
+ @Test
+ void testMapEntryConstructorWithStringValueTrimming() {
+ Map.Entry<String, Object> entry =
+ new AbstractMap.SimpleEntry<>("book_intro_vector", "
book_intro ");
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+ Assertions.assertEquals("book_intro", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ }
+
+ @Test
+ void testMapEntryConstructorWithNullKey() {
+ Map.Entry<String, Object> entry = new AbstractMap.SimpleEntry<>(null,
"book_intro");
+ IllegalArgumentException exception =
+ Assertions.assertThrows(
+ IllegalArgumentException.class, () -> new
VectorFieldSpec(entry));
+ Assertions.assertTrue(
+ exception.getMessage().contains("Field config name cannot be
null or empty"));
+ }
+
+ @Test
+ void testMapEntryConstructorWithEmpty() {
+ Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_intro_vector", null);
+ IllegalArgumentException exception =
+ Assertions.assertThrows(
+ IllegalArgumentException.class, () -> new
VectorFieldSpec(entry));
+ Assertions.assertTrue(exception.getMessage().contains("Field config
value cannot be null"));
+
+ Map.Entry<String, Object> entry2 = new
AbstractMap.SimpleEntry<>("book_intro_vector", "");
+ exception =
+ Assertions.assertThrows(
+ IllegalArgumentException.class, () -> new
VectorFieldSpec(entry2));
+ Assertions.assertTrue(
+ exception.getMessage().contains("Invalid field spec for output
field"));
+ }
+
+ @Test
+ void testMapEntryConstructorWithMapValue() {
+ Map<String, Object> fieldConfig = new HashMap<>();
+ fieldConfig.put("field", "book_image");
+ fieldConfig.put("modality", "jpeg");
+ fieldConfig.put("format", "binary");
+
+ Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_field", fieldConfig);
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+
+ Assertions.assertEquals("book_image", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.JPEG,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.BINARY,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertTrue(vectorFieldSpec.isMultimodalField());
+ Assertions.assertTrue(srcFieldSpec.isBinary());
+ }
+
+ @Test
+ void testMapEntryConstructorWithMapValueNoModality() {
+ Map<String, Object> fieldConfig = new HashMap<>();
+ fieldConfig.put("field", "book_intro");
+ fieldConfig.put("modality", "text");
+ fieldConfig.put("format", "text");
+
+ Map.Entry<String, Object> entry = new
AbstractMap.SimpleEntry<>("book_field", fieldConfig);
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+
+ Assertions.assertEquals("book_intro", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(vectorFieldSpec.isMultimodalField());
+ }
+
+ @Test
+ void testMapEntryConstructorWithInvalidListValue() {
+ List<String> textFieldConfig = Arrays.asList("text_field_1",
"text_field_2");
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "jpeg");
+ imageFieldConfig.put("format", "url");
+
+ Map.Entry<String, Object> entry =
+ new AbstractMap.SimpleEntry<>(
+ "vector_field", Arrays.asList(textFieldConfig,
imageFieldConfig));
+ IllegalArgumentException exception =
+ Assertions.assertThrows(
+ IllegalArgumentException.class, () -> new
VectorFieldSpec(entry));
+ Assertions.assertTrue(
+ exception.getMessage().contains("Invalid field spec for output
field"));
+ }
+
+ @Test
+ void testMapEntryConstructorWithSameModalityListValue() {
+ Map.Entry<String, Object> entry =
+ new AbstractMap.SimpleEntry<>(
+ "vector_field", Arrays.asList("text_field_1",
"text_field_2"));
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ Assertions.assertEquals("vector_field",
vectorFieldSpec.getFieldName());
+ Assertions.assertTrue(vectorFieldSpec.isMultimodalField());
+
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+ Assertions.assertEquals("text_field_1", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+
+ srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(1);
+ Assertions.assertEquals("text_field_2", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+ }
+
+ @Test
+ void testMapEntryConstructorWithDifferentModalityListValue() {
+ Map<String, Object> imageFieldConfig = new HashMap<>();
+ imageFieldConfig.put("field", "image_field");
+ imageFieldConfig.put("modality", "jpeg");
+ imageFieldConfig.put("format", "url");
+
+ Map<String, Object> videoFieldConfig = new HashMap<>();
+ videoFieldConfig.put("field", "video_field");
+ videoFieldConfig.put("modality", "mp4");
+ videoFieldConfig.put("format", "url");
+
+ Map.Entry<String, Object> entry =
+ new AbstractMap.SimpleEntry<>(
+ "vector_field",
+ Arrays.asList("text_field", imageFieldConfig,
videoFieldConfig));
+ VectorFieldSpec vectorFieldSpec = new VectorFieldSpec(entry);
+ Assertions.assertEquals("vector_field",
vectorFieldSpec.getFieldName());
+ Assertions.assertTrue(vectorFieldSpec.isMultimodalField());
+
+ SrcFieldSpec srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(0);
+ Assertions.assertEquals("text_field", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.TEXT,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.TEXT,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+
+ srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(1);
+ Assertions.assertEquals("image_field", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.JPEG,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.URL,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+
+ srcFieldSpec = vectorFieldSpec.getSrcFieldSpecs().get(2);
+ Assertions.assertEquals("video_field", srcFieldSpec.getFieldName());
+ Assertions.assertEquals(ModalityType.MP4,
srcFieldSpec.getModalityType());
+ Assertions.assertEquals(PayloadFormat.URL,
srcFieldSpec.getPayloadFormat());
+ Assertions.assertFalse(srcFieldSpec.isBinary());
+ }
+}