openai/openai-dotnet

Public

mirrored from https://github.com/openai/openai-dotnetAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
OpenAI_2.3.0

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

tests/Embeddings/EmbeddingsTests.cs

239lines · modecode

1using NUnit.Framework;
2using OpenAI.Embeddings;
3using OpenAI.Tests.Utility;
4using System;
5using System.ClientModel;
6using System.ClientModel.Primitives;
7using System.Collections.Generic;
8using System.Threading.Tasks;
9using static OpenAI.Tests.TestHelpers;
10
11namespace OpenAI.Tests.Embeddings;
12
13[TestFixture(true)]
14[TestFixture(false)]
15[Parallelizable(ParallelScope.All)]
16[Category("Embeddings")]
17public class EmbeddingsTests : SyncAsyncTestBase
18{
19 private static EmbeddingClient GetTestClient() => GetTestClient<EmbeddingClient>(TestScenario.Embeddings);
20
21 public EmbeddingsTests(bool isAsync) : base(isAsync)
22 {
23 }
24
25 public enum EmbeddingsInputKind
26 {
27 UsingStrings,
28 UsingIntegers,
29 }
30
31 [Test]
32 public async Task GenerateSingleEmbedding()
33 {
34 EmbeddingClient client = new("text-embedding-3-small", Environment.GetEnvironmentVariable("OPENAI_API_KEY"));
35
36 string input = "Hello, world!";
37
38 OpenAIEmbedding embedding = IsAsync
39 ? await client.GenerateEmbeddingAsync(input)
40 : client.GenerateEmbedding(input);
41 Assert.That(embedding, Is.Not.Null);
42 Assert.That(embedding.Index, Is.EqualTo(0));
43
44 ReadOnlyMemory<float> vector = embedding.ToFloats();
45 Assert.That(vector, Is.Not.Null);
46 Assert.That(vector.Span.Length, Is.EqualTo(1536));
47
48 float[] array = vector.ToArray();
49 Assert.That(array.Length, Is.EqualTo(1536));
50 }
51
52 [Test]
53 [TestCase(EmbeddingsInputKind.UsingStrings)]
54 [TestCase(EmbeddingsInputKind.UsingIntegers)]
55 public async Task GenerateMultipleEmbeddings(EmbeddingsInputKind embeddingsInputKind)
56 {
57 EmbeddingClient client = new("text-embedding-3-small", Environment.GetEnvironmentVariable("OPENAI_API_KEY"));
58
59 const int Dimensions = 456;
60
61 EmbeddingGenerationOptions options = new()
62 {
63 Dimensions = Dimensions,
64 };
65
66 OpenAIEmbeddingCollection embeddings = null;
67
68 if (embeddingsInputKind == EmbeddingsInputKind.UsingStrings)
69 {
70 List<string> prompts =
71 [
72 "Hello, world!",
73 "This is a test.",
74 "Goodbye!"
75 ];
76
77 embeddings = IsAsync
78 ? await client.GenerateEmbeddingsAsync(prompts, options)
79 : client.GenerateEmbeddings(prompts, options);
80 }
81 else if (embeddingsInputKind == EmbeddingsInputKind.UsingIntegers)
82 {
83 List<ReadOnlyMemory<int>> prompts =
84 [
85 new[] { 104, 101, 108, 108, 111 },
86 new[] { 119, 111, 114, 108, 100 },
87 new[] { 84, 69, 83, 84 }
88 ];
89
90 embeddings = IsAsync
91 ? await client.GenerateEmbeddingsAsync(prompts, options)
92 : client.GenerateEmbeddings(prompts, options);
93 }
94
95 Assert.That(embeddings, Is.Not.Null);
96 Assert.That(embeddings.Count, Is.EqualTo(3));
97 Assert.That(embeddings.Model, Is.EqualTo("text-embedding-3-small"));
98 Assert.That(embeddings.Usage.InputTokenCount, Is.GreaterThan(0));
99 Assert.That(embeddings.Usage.TotalTokenCount, Is.GreaterThan(0));
100
101 for (int i = 0; i < embeddings.Count; i++)
102 {
103 Assert.That(embeddings[i].Index, Is.EqualTo(i));
104
105 ReadOnlyMemory<float> vector = embeddings[i].ToFloats();
106 Assert.That(vector, Is.Not.Null);
107 Assert.That(vector.Span.Length, Is.EqualTo(Dimensions));
108
109 float[] array = vector.ToArray();
110 Assert.That(array.Length, Is.EqualTo(Dimensions));
111 }
112 }
113
114 [Test]
115 public async Task BadOptions()
116 {
117 EmbeddingClient client = GetTestClient();
118
119 EmbeddingGenerationOptions options = new()
120 {
121 Dimensions = -42,
122 };
123
124 Exception caughtException = null;
125
126 try
127 {
128 _ = IsAsync
129 ? await client.GenerateEmbeddingAsync("foo", options)
130 : client.GenerateEmbedding("foo", options);
131 }
132 catch (Exception ex)
133 {
134 caughtException = ex;
135 }
136
137 Assert.That(caughtException, Is.InstanceOf<ClientResultException>());
138 Assert.That(caughtException.Message, Contains.Substring("dimensions"));
139 }
140
141 [Test]
142 [TestCase(EmbeddingsInputKind.UsingStrings)]
143 [TestCase(EmbeddingsInputKind.UsingIntegers)]
144 public async Task GenerateMultipleEmbeddingsWithBadOptions(EmbeddingsInputKind embeddingsInputKind)
145 {
146 EmbeddingClient client = GetTestClient();
147
148 EmbeddingGenerationOptions options = new()
149 {
150 Dimensions = -42,
151 };
152
153 Exception caughtException = null;
154
155 try
156 {
157 if (embeddingsInputKind == EmbeddingsInputKind.UsingStrings)
158 {
159 _ = IsAsync
160 ? await client.GenerateEmbeddingsAsync(["prompt"], options)
161 : client.GenerateEmbeddings(["prompt"], options);
162 }
163 else if (embeddingsInputKind == EmbeddingsInputKind.UsingIntegers)
164 {
165 _ = IsAsync
166 ? await client.GenerateEmbeddingsAsync([new[] { 1 }], options)
167 : client.GenerateEmbeddings([new[] { 1 }], options);
168 }
169 }
170 catch (Exception ex)
171 {
172 caughtException = ex;
173 }
174
175 Assert.That(caughtException, Is.InstanceOf<ClientResultException>());
176 Assert.That(caughtException.Message, Contains.Substring("dimensions"));
177 }
178
179 [Test]
180 public async Task EmbeddingFromStringAndEmbeddingFromTokenIdsAreEqual()
181 {
182 EmbeddingClient client = GetTestClient();
183
184 List<string> input1 = new List<string> { "Hello, world!" };
185 List<ReadOnlyMemory<int>> input2 = new List<ReadOnlyMemory<int>> { new[] { 9906, 11, 1917, 0 } };
186
187 OpenAIEmbeddingCollection results1 = IsAsync
188 ? await client.GenerateEmbeddingsAsync(input1)
189 : client.GenerateEmbeddings(input1);
190
191 OpenAIEmbeddingCollection results2 = IsAsync
192 ? await client.GenerateEmbeddingsAsync(input2)
193 : client.GenerateEmbeddings(input2);
194
195 ReadOnlyMemory<float> vector1 = results1[0].ToFloats();
196 ReadOnlyMemory<float> vector2 = results2[0].ToFloats();
197
198 for (int i = 0; i < vector1.Length; i++)
199 {
200 Assert.That(vector1.Span[i], Is.EqualTo(vector2.Span[i]).Within(0.0005));
201 }
202 }
203
204 [Test]
205 public void SerializeEmbeddingCollection()
206 {
207 // TODO: Add this test.
208 }
209
210 [Test]
211 public void JsonArraySupport()
212 {
213 string json = """
214 {
215 "object":"list",
216 "data":[
217 {
218 "object":"embedding",
219 "embedding":[-0.011229509,0.107915245,-0.15163477]
220 }
221 ]
222 }
223 """;
224
225 BinaryData binaryData = BinaryData.FromString(json);
226
227 OpenAIEmbeddingCollection embeddings = ModelReaderWriter.Read<OpenAIEmbeddingCollection>(binaryData);
228
229 Assert.That(embeddings, Is.Not.Null);
230 Assert.That(embeddings.Count, Is.EqualTo(1));
231 var embedding = embeddings[0];
232 Assert.That(embedding, Is.Not.Null);
233 ReadOnlySpan<float> vector = embedding.ToFloats().Span;
234 Assert.That(vector.Length, Is.EqualTo(3));
235 Assert.That(vector[0], Is.EqualTo(-0.011229509f));
236 Assert.That(vector[1], Is.EqualTo(0.107915245f));
237 Assert.That(vector[2], Is.EqualTo(-0.15163477f));
238 }
239}
240