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
5 changes: 3 additions & 2 deletions docs/release_notes.md
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,12 @@

### 🔧 Compatibility Notes

- [Foundation models] SAP-RPT was updated to the newer 1.6.0 API
- [RPT] SAP-RPT was updated to the newer 1.6.0 API

### ✨ New Functionality

- [Grounding] Enable presigned URL for Pipeline Documents.
- [Grounding] Enabled presigned URL for Pipeline Documents.
- [RPT] Added `SAP_RPT_1_6` and `SAP_RPT_1_6_LARGE` to the model list in `RptModel`.

### 📈 Improvements

Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package com.sap.ai.sdk.foundationmodels.rpt;

import static com.sap.ai.sdk.core.JacksonConfiguration.getDefaultObjectMapper;
import static com.sap.ai.sdk.foundationmodels.rpt.RptModel.SAP_RPT_1_6_LARGE;

import com.fasterxml.jackson.core.JsonProcessingException;
import com.google.common.annotations.Beta;
Expand Down Expand Up @@ -31,15 +32,9 @@
public class RptClient {
@Nonnull private final DefaultApi api;
@Nonnull private final DefaultApi apiWithGzipEncoding;
private final boolean usesOldModel;
private final boolean contextModePossible;

private static final Set<String> PRE1_6MODELS =
Set.of(
"sap-rpt-1-large",
"sap-rpt-1-small",
"sap-rpt-1.1-preview",
"sap-rpt-1.5",
"sap-rpt-1.5-large");
private static final Set<RptModel> MODELS_WITH_CONTEXT_MODE = Set.of(SAP_RPT_1_6_LARGE);

/**
* Creates a new RptClient for the specified foundation model.
Expand All @@ -51,9 +46,9 @@ public class RptClient {
@Nonnull
public static RptClient forModel(@Nonnull final RptModel foundationModel)
throws DeploymentResolutionException {
final var usesOldModel = PRE1_6MODELS.contains(foundationModel.name());
final var contextModePossible = MODELS_WITH_CONTEXT_MODE.contains(foundationModel);
final var destination = new AiCoreService().getInferenceDestination().forModel(foundationModel);
return forDestination(destination, usesOldModel);
return forDestination(destination, contextModePossible);
}

/**
Expand All @@ -63,11 +58,11 @@ public static RptClient forModel(@Nonnull final RptModel foundationModel)
* @return A new instance of RptClient.
*/
static RptClient forDestination(
@Nonnull final Destination destination, final boolean usesOldModel) {
@Nonnull final Destination destination, final boolean contextModePossible) {
final var apiClient = ApiClient.create(destination).withObjectMapper(getDefaultObjectMapper());
final var api = new DefaultApi(apiClient);
return new RptClient(
api, api.withDefaultHeaders(Map.of("Content-Encoding", "gzip")), usesOldModel);
api, api.withDefaultHeaders(Map.of("Content-Encoding", "gzip")), contextModePossible);
}

/**
Expand All @@ -88,15 +83,14 @@ static RptClient forDestination(
*
* @param requestBody The prediction request
* @return prediction response from the RPT model
* @apiNote When used with a pre-1.6 model, the {@code contextMode} field of the embedded {@link
* com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is set to {@code
* null} on the passed-in object as a side effect.
* @apiNote When used with a model that does not support it, the {@code contextMode} field of the
* embedded {@link com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is
* set to {@code null} on the passed-in object as a side effect.
*/
@Beta
@Nonnull
public PredictResponsePayload tableCompletion(@Nonnull final PredictRequestPayload requestBody) {
// contextMode has to be null for models < 1.6
if (usesOldModel) {
if (!contextModePossible) {
configFrom(requestBody).setContextMode(null);
}
return apiWithGzipEncoding.predict(requestBody);
Expand Down Expand Up @@ -133,17 +127,16 @@ private static PredictionConfig configFrom(@Nonnull final PredictRequestPayload
* @param parquetFile Parquet file
* @param predictionConfig The prediction configuration
* @return prediction response from the RPT model
* @apiNote When used with a pre-1.6 model, the {@code contextMode} field of the passed-in {@link
* com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is set to {@code
* null} as a side effect.
* @apiNote When used with a model that does not support it, the {@code contextMode} field of the
* passed-in {@link com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is
* set to {@code null} as a side effect.
* @since 1.16.0
*/
@Beta
@Nonnull
public PredictResponsePayload tableCompletion(
@Nonnull final File parquetFile, @Nonnull final PredictionConfig predictionConfig) {
// contextMode has to be null for models < 1.6
if (usesOldModel) {
if (!contextModePossible) {
predictionConfig.setContextMode(null);
}
try {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,12 @@ public record RptModel(@Nonnull String name, @Nullable String version) implement
/** SAP Relational Pre-trained Transformer 1.5 Large model. */
public static final RptModel SAP_RPT_1_5_LARGE = new RptModel("sap-rpt-1.5-large", null);

/** SAP Relational Pre-trained Transformer 1.6 model. */
public static final RptModel SAP_RPT_1_6 = new RptModel("sap-rpt-1.6", null);

/** SAP Relational Pre-trained Transformer 1.6 Large model. */
public static final RptModel SAP_RPT_1_6_LARGE = new RptModel("sap-rpt-1.6-large", null);
Comment thread
CharlesDuboisSAP marked this conversation as resolved.

/**
* Create a new instance of RptModel with the provided version.
*
Expand Down
Original file line number Diff line number Diff line change
@@ -1,9 +1,13 @@
package com.sap.ai.sdk.foundationmodels.rpt;

import static com.github.tomakehurst.wiremock.client.WireMock.aResponse;
import static com.github.tomakehurst.wiremock.client.WireMock.containing;
import static com.github.tomakehurst.wiremock.client.WireMock.not;
import static com.github.tomakehurst.wiremock.client.WireMock.post;
import static com.github.tomakehurst.wiremock.client.WireMock.postRequestedFor;
import static com.github.tomakehurst.wiremock.client.WireMock.stubFor;
import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo;
import static com.github.tomakehurst.wiremock.client.WireMock.verify;
import static com.sap.ai.sdk.foundationmodels.rpt.generated.model.ColumnType.STRING;
import static com.sap.ai.sdk.foundationmodels.rpt.generated.model.TargetColumnConfig.TaskTypeEnum.CLASSIFICATION;
import static org.assertj.core.api.Assertions.assertThat;
Expand Down Expand Up @@ -64,7 +68,7 @@ class RptClientTest {
void setup(final WireMockRuntimeInfo server) {
final DefaultHttpDestination destination =
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build();
client = RptClient.forDestination(destination, false);
client = RptClient.forDestination(destination, true);
ApacheHttpClient5Accessor.setHttpClientCache(ApacheHttpClient5Cache.DISABLED);
}

Expand Down Expand Up @@ -310,7 +314,7 @@ void testTableCompletionWithParquetThrowsIllegalArgumentException() {
@Test
void testOldModelThrowsOnUnknownPayloadType() {
final var oldModelClient =
RptClient.forDestination(DefaultHttpDestination.builder("http://localhost").build(), true);
RptClient.forDestination(DefaultHttpDestination.builder("http://localhost").build(), false);
final var unknownPayload = mock(PredictRequestPayload.class);

assertThatThrownBy(() -> oldModelClient.tableCompletion(unknownPayload))
Expand All @@ -319,11 +323,11 @@ void testOldModelThrowsOnUnknownPayloadType() {
}

@Test
void testOldModelStripsContextModeFromRowWiseRequest(final WireMockRuntimeInfo server) {
void testContextModeStrippedFromRowWiseRequest(final WireMockRuntimeInfo server) {
stubFor(post(urlEqualTo("/predict")).willReturn(aResponse().withStatus(200).withBody("{}")));
final var oldModelClient =
RptClient.forDestination(
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true);
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), false);

val config =
PredictionConfig.create()
Expand All @@ -340,14 +344,16 @@ void testOldModelStripsContextModeFromRowWiseRequest(final WireMockRuntimeInfo s
assertThat(config.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEEP);
oldModelClient.tableCompletion(request);
assertThat(config.getContextMode()).isNull();
verify(
postRequestedFor(urlEqualTo("/predict")).withRequestBody(not(containing("context_mode"))));
}

@Test
void testOldModelStripsContextModeFromColumnWiseRequest(final WireMockRuntimeInfo server) {
void testContextModeStrippedFromColumnWiseRequest(final WireMockRuntimeInfo server) {
stubFor(post(urlEqualTo("/predict")).willReturn(aResponse().withStatus(200).withBody("{}")));
final var oldModelClient =
RptClient.forDestination(
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true);
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), false);

val config =
PredictionConfig.create()
Expand All @@ -364,6 +370,8 @@ void testOldModelStripsContextModeFromColumnWiseRequest(final WireMockRuntimeInf
assertThat(config.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEFAULT);
oldModelClient.tableCompletion(request);
assertThat(config.getContextMode()).isNull();
verify(
postRequestedFor(urlEqualTo("/predict")).withRequestBody(not(containing("context_mode"))));
}

@Test
Expand All @@ -373,7 +381,7 @@ void testOldModelStripsContextModeFromParquetRequest(final WireMockRuntimeInfo s
.willReturn(aResponse().withStatus(200).withBody("{}")));
final var oldModelClient =
RptClient.forDestination(
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true);
DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), false);

val parquetFile = Path.of("src/test/resources/rpt/test-data.parquet").toFile();
val predictionConfig =
Expand All @@ -384,5 +392,8 @@ void testOldModelStripsContextModeFromParquetRequest(final WireMockRuntimeInfo s
assertThat(predictionConfig.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEEP);
oldModelClient.tableCompletion(parquetFile, predictionConfig);
assertThat(predictionConfig.getContextMode()).isNull();
verify(
postRequestedFor(urlEqualTo("/predict_parquet"))
.withRequestBody(not(containing("context_mode"))));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import com.sap.ai.sdk.core.model.AiModelBaseData;
import com.sap.ai.sdk.core.model.AiModelVersion;
import com.sap.ai.sdk.foundationmodels.openai.OpenAiModel;
import com.sap.ai.sdk.foundationmodels.rpt.RptModel;
import com.sap.ai.sdk.orchestration.OrchestrationAiModel;
import java.lang.reflect.Field;
import java.util.HashMap;
Expand Down Expand Up @@ -141,6 +142,53 @@ void orchestrationAiModelAvailability() {
softly.assertAll();
}

@Test
@DisplayName("Declared RPT models must be superset of our AI Core account's available RPT models")
@SneakyThrows
void rptModelAvailability() {

// Gather AI Core's list of available RPT models
val aiModelList = new ScenarioController().getModels().getResources();

val internalOnlyModels = Set.of("sap-rpt-1.1-preview");

val availableRptModels =
aiModelList.stream()
.filter(model -> model.getModel().contains("rpt"))
.filter(model -> !internalOnlyModels.contains(model.getModel()))
.collect(
() -> new HashMap<String, Boolean>(),
(list, model) -> list.put(model.getModel(), isDeprecated(model)),
HashMap::putAll);

// Gather our declared RPT models
Field[] declaredFields = RptModel.class.getFields();

// get the models from the OpenAiModel class
HashMap<String, Boolean> declaredRptModelList = new HashMap<>();
for (Field field : declaredFields) {
if (field.getType().equals(RptModel.class)) {
declaredRptModelList.put(
((RptModel) field.get(null)).name(), field.isAnnotationPresent(Deprecated.class));
}
}

// Assert that the declared RPT models match the expected list
assertThat(declaredRptModelList.keySet()).containsAll(availableRptModels.keySet());

SoftAssertions softly = new SoftAssertions();
for (val model : availableRptModels.entrySet()) {
Boolean declaredDeprecated = declaredRptModelList.get(model.getKey());
softly
.assertThat(declaredDeprecated)
.withFailMessage(
"%s is deprecated:%s on AI Core but deprecated:%s in AI SDK",
model.getKey(), model.getValue(), declaredDeprecated)
.isEqualTo(model.getValue());
}
softly.assertAll();
}

private static boolean isDeprecated(AiModelBaseData model) {
Optional<AiModelVersion> version =
model.getVersions().stream().filter(AiModelVersion::isIsLatest).findFirst();
Expand Down