mirror of
https://github.com/MindWorkAI/AI-Studio.git
synced 2026-10-05 06:49:40 +00:00
made tokenizer concurrent and threadsafe
This commit is contained in:
1 parent
6985678d2f
commit
c8c1e38475
11 files changed
+110
-306
No files matched your search
@@ -1259,12 +1259,12 @@ public partial class ChatComponent : MSGComponentBase, IAsyncDisposable
|
||||
return;
|
||||
}
|
||||
|
||||
var response = await this.RustService.GetTokenCount(this.Provider.InstanceName, this.Provider.TokenizerPath, this.inputField.Value);
|
||||
var response = await this.RustService.GetTokenCount(this.Provider, this.inputField.Value);
|
||||
if (response is null)
|
||||
return;
|
||||
if (!response.Value.Success)
|
||||
{
|
||||
this.Logger.LogWarning($"Failed to calculate token count: status='{response.Value.Status}', reason='{response.Value.Message}'");
|
||||
this.Logger.LogWarning("Failed to calculate token count: reason='{Reason}'", response.Value.Message);
|
||||
return;
|
||||
}
|
||||
this.tokenCount = response.Value.TokenCount.ToString();
|
||||
|
||||
@@ -4,5 +4,4 @@ public readonly record struct TokenizerResponse(
|
||||
bool Success,
|
||||
int TokenCount,
|
||||
string Message,
|
||||
TokenizerStatus Status = TokenizerStatus.UNAVAILABLE,
|
||||
string StoredPath = "");
|
||||
@@ -1,8 +0,0 @@
|
||||
namespace AIStudio.Tools.Rust;
|
||||
|
||||
public enum TokenizerStatus
|
||||
{
|
||||
UNAVAILABLE,
|
||||
RUNNING,
|
||||
AVAILABLE,
|
||||
}
|
||||
@@ -66,7 +66,7 @@ public sealed partial class DataSourceEmbeddingService
|
||||
{
|
||||
var segments = new List<ExtractedFileSegment>();
|
||||
|
||||
await foreach (var segment in rustService.StreamArbitraryFileDataWithTokenCounts(filePath, embeddingProvider.Name, embeddingProvider.TokenizerPath, token))
|
||||
await foreach (var segment in rustService.StreamArbitraryFileDataWithTokenCounts(filePath, embeddingProvider, token))
|
||||
{
|
||||
var normalized = NormalizeChunkSegment(segment.Content);
|
||||
if (!string.IsNullOrWhiteSpace(normalized))
|
||||
@@ -536,8 +536,8 @@ public sealed partial class DataSourceEmbeddingService
|
||||
|
||||
private async Task<int> GetEmbeddingTokenCountAsync(EmbeddingProvider embeddingProvider, string text, CancellationToken token)
|
||||
{
|
||||
var response = await rustService.GetTokenCount(embeddingProvider.Name, embeddingProvider.TokenizerPath, text, token);
|
||||
if (response is { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
var response = await rustService.GetTokenCount(embeddingProvider, text, token);
|
||||
if (response is { Success: true })
|
||||
return response.Value.TokenCount;
|
||||
|
||||
var message = response?.Message ?? "No response was returned by the tokenizer service.";
|
||||
|
||||
@@ -6,7 +6,6 @@ using AIStudio.Tools.Databases;
|
||||
using AIStudio.Tools.Databases.EmbeddingState;
|
||||
using AIStudio.Tools.Databases.VectorStore;
|
||||
using AIStudio.Tools.RAG;
|
||||
using AIStudio.Tools.Rust;
|
||||
|
||||
namespace AIStudio.Tools.Services;
|
||||
|
||||
@@ -160,12 +159,8 @@ public sealed class DataSourceLocalRetrievalService(
|
||||
return false;
|
||||
}
|
||||
|
||||
var tokenCountResponse = await rustService.GetTokenCount(
|
||||
embeddingProvider.Name,
|
||||
embeddingProvider.TokenizerPath,
|
||||
query,
|
||||
token);
|
||||
if (tokenCountResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
var tokenCountResponse = await rustService.GetTokenCount(embeddingProvider, query, token);
|
||||
if (tokenCountResponse is not { Success: true })
|
||||
{
|
||||
logger.LogWarning(
|
||||
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the token count for embedding provider '{EmbeddingProviderName}' could not be determined. Reason='{Reason}'.",
|
||||
|
||||
@@ -2,6 +2,7 @@ using System.Text;
|
||||
using System.Text.Json;
|
||||
using System.Runtime.CompilerServices;
|
||||
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.Rust;
|
||||
|
||||
namespace AIStudio.Tools.Services;
|
||||
@@ -85,37 +86,21 @@ public sealed partial class RustService
|
||||
|
||||
public async IAsyncEnumerable<string> StreamArbitraryFileData(string path, bool extractImages = false, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
await foreach (var segment in this.StreamArbitraryFileDataCore(path, extractImages, false, token))
|
||||
await foreach (var segment in this.StreamArbitraryFileDataCore(path, extractImages, false, string.Empty, token))
|
||||
yield return segment.Content;
|
||||
}
|
||||
|
||||
public async IAsyncEnumerable<ArbitraryFileDataSegment> StreamArbitraryFileDataWithTokenCounts(
|
||||
string path,
|
||||
string providerName,
|
||||
string tokenizerPath,
|
||||
EmbeddingProvider embeddingProvider,
|
||||
[EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
await this.tokenizerLock.WaitAsync(token);
|
||||
try
|
||||
await foreach (var segment in this.StreamArbitraryFileDataCore(path, false, true, embeddingProvider.TokenizerPath, token))
|
||||
{
|
||||
var tokenizerResponse = await this.EnsureTokenizerCoreAsync(providerName, tokenizerPath);
|
||||
if (tokenizerResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
{
|
||||
var message = tokenizerResponse?.Message ?? "No response was returned by the tokenizer service.";
|
||||
throw new InvalidOperationException($"Could not initialize tokenizer for provider '{providerName}'. {message}");
|
||||
}
|
||||
if (segment.TokenCount is null)
|
||||
throw new InvalidOperationException($"Rust did not return a token count for an extracted segment from '{path}' using provider '{embeddingProvider.Name}'.");
|
||||
|
||||
await foreach (var segment in this.StreamArbitraryFileDataCore(path, false, true, token))
|
||||
{
|
||||
if (segment.TokenCount is null)
|
||||
throw new InvalidOperationException($"Rust did not return a token count for an extracted segment from '{path}'.");
|
||||
|
||||
yield return new(segment.Content, segment.TokenCount.Value);
|
||||
}
|
||||
}
|
||||
finally
|
||||
{
|
||||
this.tokenizerLock.Release();
|
||||
yield return new(segment.Content, segment.TokenCount.Value);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -123,10 +108,11 @@ public sealed partial class RustService
|
||||
string path,
|
||||
bool extractImages,
|
||||
bool includeTokenCount,
|
||||
string tokenizerPath,
|
||||
[EnumeratorCancellation] CancellationToken token)
|
||||
{
|
||||
var streamId = Guid.NewGuid().ToString();
|
||||
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count={includeTokenCount}";
|
||||
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count={includeTokenCount}&tokenizer_path={Uri.EscapeDataString(tokenizerPath)}";
|
||||
using var request = new HttpRequestMessage(HttpMethod.Get, requestUri);
|
||||
using var response = await this.http.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, token);
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
using AIStudio.Provider;
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.Rust;
|
||||
|
||||
namespace AIStudio.Tools.Services;
|
||||
@@ -7,35 +7,12 @@ public sealed partial class RustService
|
||||
{
|
||||
internal const int MAX_TOKEN_COUNT_REQUEST_TEXT_LENGTH = 200_000;
|
||||
|
||||
private readonly SemaphoreSlim tokenizerLock = new(1, 1);
|
||||
private string currentTokenizerPath = string.Empty;
|
||||
private bool hasInitializedTokenizer;
|
||||
|
||||
private static TokenizerResponse CreateUnavailableTokenizerResponse(string message) => new(
|
||||
false,
|
||||
0,
|
||||
message,
|
||||
TokenizerStatus.UNAVAILABLE,
|
||||
string.Empty);
|
||||
|
||||
public async Task<TokenizerResponse> GetTokenizerInfo(CancellationToken cancellationToken = default)
|
||||
{
|
||||
try
|
||||
{
|
||||
return await this.http.GetFromJsonAsync<TokenizerResponse>("/system/tokenizer/info", this.jsonRustSerializerOptions, cancellationToken);
|
||||
}
|
||||
catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested)
|
||||
{
|
||||
this.logger?.LogWarning("Fetching tokenizer info from Rust service was cancelled by caller.");
|
||||
return CreateUnavailableTokenizerResponse("Operation cancelled by caller.");
|
||||
}
|
||||
catch (Exception e)
|
||||
{
|
||||
this.logger?.LogError(e, "Error while fetching tokenizer info from Rust service.");
|
||||
return CreateUnavailableTokenizerResponse(e.Message);
|
||||
}
|
||||
}
|
||||
|
||||
public async Task<TokenizerResponse> ValidateTokenizer(string filePath)
|
||||
{
|
||||
var result = await this.http.PostAsJsonAsync("/tokenizer/validate", new {
|
||||
@@ -86,99 +63,25 @@ public sealed partial class RustService
|
||||
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
|
||||
}
|
||||
|
||||
public Task<TokenizerResponse?> GetTokenCount(string text)
|
||||
{
|
||||
return this.GetTokenCountCoreAsync(text);
|
||||
}
|
||||
public Task<TokenizerResponse?> GetTokenCount(AIStudio.Settings.Provider provider, string text, CancellationToken cancellationToken = default) =>
|
||||
this.GetTokenCount(provider.InstanceName, provider.TokenizerPath, text, cancellationToken);
|
||||
|
||||
public async Task<TokenizerResponse?> GetTokenCount(string providerName, string path, string text, CancellationToken cancellationToken = default)
|
||||
{
|
||||
await this.tokenizerLock.WaitAsync(cancellationToken);
|
||||
try
|
||||
{
|
||||
var tokenizerResponse = await this.EnsureTokenizerCoreAsync(providerName, path);
|
||||
if (tokenizerResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
return tokenizerResponse;
|
||||
public Task<TokenizerResponse?> GetTokenCount(EmbeddingProvider provider, string text, CancellationToken cancellationToken = default) =>
|
||||
this.GetTokenCount(provider.Name, provider.TokenizerPath, text, cancellationToken);
|
||||
|
||||
return await this.GetTokenCountCoreAsync(text, cancellationToken);
|
||||
}
|
||||
finally
|
||||
{
|
||||
this.tokenizerLock.Release();
|
||||
}
|
||||
}
|
||||
|
||||
private async Task<TokenizerResponse?> GetTokenCountCoreAsync(string text, CancellationToken cancellationToken = default)
|
||||
private async Task<TokenizerResponse?> GetTokenCount(string providerName, string tokenizerPath, string text, CancellationToken cancellationToken)
|
||||
{
|
||||
var result = await this.http.PostAsJsonAsync("/tokenizer/count", new {
|
||||
text = text,
|
||||
tokenizer_path = tokenizerPath,
|
||||
}, this.jsonRustSerializerOptions, cancellationToken);
|
||||
|
||||
if (!result.IsSuccessStatusCode)
|
||||
{
|
||||
this.logger!.LogError($"Failed to get the token count '{result.StatusCode}'");
|
||||
this.hasInitializedTokenizer = false;
|
||||
this.logger!.LogError("Failed to get the token count for provider '{ProviderName}': {StatusCode}", providerName, result.StatusCode);
|
||||
return CreateUnavailableTokenizerResponse("Error while getting token count from Rust service: "+result.StatusCode);
|
||||
}
|
||||
|
||||
var response = await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
|
||||
if (response is not { Status: TokenizerStatus.AVAILABLE })
|
||||
this.hasInitializedTokenizer = false;
|
||||
|
||||
return response;
|
||||
}
|
||||
|
||||
public async Task<TokenizerResponse?> SetTokenizer(string providerName, string path)
|
||||
{
|
||||
this.logger!.LogInformation($"Setting a new tokenizer for '{providerName}'");
|
||||
var result = await this.http.PostAsJsonAsync("/tokenizer/set", new {
|
||||
file_path = path,
|
||||
}, this.jsonRustSerializerOptions);
|
||||
|
||||
if (!result.IsSuccessStatusCode)
|
||||
{
|
||||
this.logger!.LogError($"Failed to set the tokenizer '{result.StatusCode}'");
|
||||
this.hasInitializedTokenizer = false;
|
||||
return CreateUnavailableTokenizerResponse("An error occured while sending the path to the Rust framework for setting a tokenizer: "+result.StatusCode);
|
||||
}
|
||||
|
||||
var response = await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
|
||||
if (response is not { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
this.hasInitializedTokenizer = false;
|
||||
|
||||
return response;
|
||||
}
|
||||
|
||||
public async Task<TokenizerResponse?> EnsureTokenizer(string providerName, string path)
|
||||
{
|
||||
await this.tokenizerLock.WaitAsync();
|
||||
try
|
||||
{
|
||||
return await this.EnsureTokenizerCoreAsync(providerName, path);
|
||||
}
|
||||
finally
|
||||
{
|
||||
this.tokenizerLock.Release();
|
||||
}
|
||||
}
|
||||
|
||||
private async Task<TokenizerResponse?> EnsureTokenizerCoreAsync(string providerName, string path)
|
||||
{
|
||||
if (this.hasInitializedTokenizer && this.currentTokenizerPath == path)
|
||||
return new TokenizerResponse(true, 0, string.Empty, TokenizerStatus.AVAILABLE);
|
||||
|
||||
var response = await this.SetTokenizer(providerName, path);
|
||||
if (response is { Success: true, Status: TokenizerStatus.AVAILABLE })
|
||||
{
|
||||
this.currentTokenizerPath = path;
|
||||
this.hasInitializedTokenizer = true;
|
||||
}
|
||||
else
|
||||
{
|
||||
this.currentTokenizerPath = string.Empty;
|
||||
this.hasInitializedTokenizer = false;
|
||||
}
|
||||
|
||||
return response;
|
||||
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions, cancellationToken);
|
||||
}
|
||||
}
|
||||
@@ -94,7 +94,6 @@ public sealed partial class RustService : BackgroundService
|
||||
{
|
||||
this.http.Dispose();
|
||||
this.userLanguageLock.Dispose();
|
||||
this.tokenizerLock.Dispose();
|
||||
this.userNameLock.Dispose();
|
||||
base.Dispose();
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user