Added a retrieval context validation agent (#289)

This commit is contained in:
Thorsten Sommer authored and GitHub committed 2025-02-22 20:51:06 +01:00
1 parent 96e6372fcd
commit bfc9f2ea1d
26 files changed
+829 -135

No files matched your search

@@ -0,0 +1,16 @@
namespace AIStudio.Tools;
/// <summary>
/// A contract for data classes with a confidence value.
/// </summary>
/// <remarks>
/// Using this confidence contract allows us to provide
/// algorithms based on confidence values.
/// </remarks>
public interface IConfidence
{
/// <summary>
/// How confident is the AI in this task or decision?
/// </summary>
public float Confidence { get; init; }
}
@@ -0,0 +1,101 @@
namespace AIStudio.Tools;
public static class IConfidenceExtensions
{
public static TargetWindow DetermineTargetWindow<T>(this IReadOnlyList<T> items, TargetWindowStrategy strategy, int numMaximumItems = 30) where T : IConfidence
{
switch (strategy)
{
case TargetWindowStrategy.A_FEW_GOOD_ONES:
return new(1, 2, 3, 0f);
case TargetWindowStrategy.TOP10_BETTER_THAN_GUESSING:
var numItemsBetterThanGuessing = items.Count(x => x.Confidence > 0.5f);
if(numItemsBetterThanGuessing < 3)
return new(1, 2, 3, 0.5f);
// We want the top 10% of items better than guessing:
var numTop10Percent = (int) MathF.Floor(numItemsBetterThanGuessing * 0.1f);
// When these 10% are just a few items, we take them all:
if (numTop10Percent <= 10)
{
var diff = numItemsBetterThanGuessing - numTop10Percent;
var num50Percent = (int) MathF.Floor(numItemsBetterThanGuessing * 0.5f);
return new(num50Percent, num50Percent + 1, Math.Max(numItemsBetterThanGuessing, diff), 0.5f);
}
// Let's define the size of the window:
const int MIN_NUM_ITEMS = 3;
var windowMin = Math.Max(MIN_NUM_ITEMS + 1, numTop10Percent);
windowMin = Math.Min(windowMin, numMaximumItems - 1);
var totalMin = Math.Max(MIN_NUM_ITEMS, windowMin - 3);
var windowSize = (int)MathF.Max(MathF.Floor(numTop10Percent * 0.1f), MathF.Min(10, numTop10Percent));
var windowMax = Math.Min(numMaximumItems, numTop10Percent + windowSize);
return new(totalMin, windowMin, windowMax, 0.5f);
case TargetWindowStrategy.NONE:
default:
return new(-1, -1, -1, 0f);
}
}
/// <summary>
/// Determine the optimal confidence threshold for a list of items
/// in order to match a target window of number of items.
/// </summary>
/// <param name="items">The list of confidence items to analyze.</param>
/// <param name="targetWindow">The target window for the number of items.</param>
/// <param name="maxSteps">The maximum number of steps to search for the threshold.</param>
/// <typeparam name="T">The type of items in the list.</typeparam>
/// <returns>The confidence threshold.</returns>
public static float GetConfidenceThreshold<T>(this IReadOnlyList<T> items, TargetWindow targetWindow, int maxSteps = 10) where T : IConfidence
{
if(!targetWindow.IsValid())
{
var logger = Program.SERVICE_PROVIDER.GetService<ILogger<IConfidence>>()!;
logger.LogWarning("The target window is invalid. Returning 0f as threshold.");
return 0f;
}
var confidenceValues = items.Select(x => x.Confidence).ToList();
var minConfidence = confidenceValues.Min();
var lowerBound = MathF.Max(minConfidence, targetWindow.MinThreshold);
var upperBound = confidenceValues.Max();
//
// We search for a threshold so that we have between
// targetWindowMin and targetWindowMax items. When not
// possible, we take all items (e.g., threshold = 0f; depends on the used window strategy)
//
var threshold = 0.0f;
// Check the case where the confidence values are too close:
if (upperBound - minConfidence >= 0.01)
{
var previousThreshold = threshold;
for (var i = 0; i < maxSteps; i++)
{
threshold = lowerBound + (upperBound - lowerBound) * i / maxSteps;
var numMatches = items.Count(x => x.Confidence >= threshold);
if (numMatches <= targetWindow.NumMinItems)
{
threshold = previousThreshold;
break;
}
if (targetWindow.InsideWindow(numMatches))
break;
previousThreshold = threshold;
}
}
else
{
var logger = Program.SERVICE_PROVIDER.GetService<ILogger<IConfidence>>()!;
logger.LogWarning("The confidence values are too close. Returning 0f as threshold.");
}
return threshold;
}
}
@@ -1,7 +1,9 @@
using System.Text;
using AIStudio.Agents;
using AIStudio.Chat;
using AIStudio.Provider;
using AIStudio.Settings;
namespace AIStudio.Tools.RAG.AugmentationProcesses;
@@ -22,13 +24,36 @@ public sealed class AugmentationOne : IAugmentationProcess
public async Task<ChatThread> ProcessAsync(IProvider provider, IContent lastPrompt, ChatThread chatThread, IReadOnlyList<IRetrievalContext> retrievalContexts, CancellationToken token = default)
{
var logger = Program.SERVICE_PROVIDER.GetService<ILogger<AugmentationOne>>()!;
var settings = Program.SERVICE_PROVIDER.GetService<SettingsManager>()!;
if(retrievalContexts.Count == 0)
{
logger.LogWarning("No retrieval contexts were issued. Skipping the augmentation process.");
return chatThread;
}
var numTotalRetrievalContexts = retrievalContexts.Count;
// Want the user to validate all retrieval contexts?
if (settings.ConfigurationData.AgentRetrievalContextValidation.EnableRetrievalContextValidation)
{
// Let's get the validation agent & set up its provider:
var validationAgent = Program.SERVICE_PROVIDER.GetService<AgentRetrievalContextValidation>()!;
validationAgent.SetLLMProvider(provider);
// Let's validate all retrieval contexts:
var validationResults = await validationAgent.ValidateRetrievalContextsAsync(lastPrompt, chatThread, retrievalContexts, token);
//
// Now, filter the retrieval contexts to the most relevant ones:
//
var targetWindow = validationResults.DetermineTargetWindow(TargetWindowStrategy.TOP10_BETTER_THAN_GUESSING);
var threshold = validationResults.GetConfidenceThreshold(targetWindow);
// Filter the retrieval contexts:
retrievalContexts = validationResults.Where(x => x.RetrievalContext is not null && x.Confidence >= threshold).Select(x => x.RetrievalContext!).ToList();
}
logger.LogInformation($"Starting the augmentation process over {numTotalRetrievalContexts:###,###,###,###} retrieval contexts.");
//
@@ -38,63 +63,8 @@ public sealed class AugmentationOne : IAugmentationProcess
sb.AppendLine("The following useful information will help you in processing the user prompt:");
sb.AppendLine();
var index = 0;
foreach(var retrievalContext in retrievalContexts)
{
index++;
sb.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
sb.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
sb.AppendLine($"Content category: {retrievalContext.Category}");
sb.AppendLine($"Content type: {retrievalContext.Type}");
sb.AppendLine($"Content path: {retrievalContext.Path}");
if(retrievalContext.Links.Count > 0)
{
sb.AppendLine("Additional links:");
foreach(var link in retrievalContext.Links)
sb.AppendLine($"- {link}");
}
switch(retrievalContext)
{
case RetrievalTextContext textContext:
sb.AppendLine();
sb.AppendLine("Matched text content:");
sb.AppendLine("````");
sb.AppendLine(textContext.MatchedText);
sb.AppendLine("````");
if(textContext.SurroundingContent.Count > 0)
{
sb.AppendLine();
sb.AppendLine("Surrounding text content:");
foreach(var surrounding in textContext.SurroundingContent)
{
sb.AppendLine();
sb.AppendLine("````");
sb.AppendLine(surrounding);
sb.AppendLine("````");
}
}
break;
case RetrievalImageContext imageContext:
sb.AppendLine();
sb.AppendLine("Matched image content as base64-encoded data:");
sb.AppendLine("````");
sb.AppendLine(await imageContext.AsBase64(token));
sb.AppendLine("````");
break;
default:
logger.LogWarning($"The retrieval content type '{retrievalContext.Type}' of data source '{retrievalContext.DataSourceName}' at location '{retrievalContext.Path}' is not supported yet.");
break;
}
sb.AppendLine();
}
// Let's convert all retrieval contexts to Markdown:
await retrievalContexts.AsMarkdown(sb, token);
//
// Append the entire augmentation to the chat thread,
@@ -70,41 +70,9 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
if (aiSelectedDataSources.Count > 3)
{
//
// We have more than 3 data sources. Let's filter by confidence.
// In order to do that, we must identify the lower and upper
// bounds of the confidence interval:
//
var confidenceValues = aiSelectedDataSources.Select(x => x.Confidence).ToList();
var lowerBound = confidenceValues.Min();
var upperBound = confidenceValues.Max();
//
// Next, we search for a threshold so that we have between 2 and 3
// data sources. When not possible, we take all data sources.
//
var threshold = 0.0f;
// Check the case where the confidence values are too close:
if (upperBound - lowerBound >= 0.01)
{
var previousThreshold = 0.0f;
for (var i = 0; i < 10; i++)
{
threshold = lowerBound + (upperBound - lowerBound) * i / 10;
var numMatches = aiSelectedDataSources.Count(x => x.Confidence >= threshold);
if (numMatches <= 1)
{
threshold = previousThreshold;
break;
}
if (numMatches is <= 3 and >= 2)
break;
previousThreshold = threshold;
}
}
// We have more than 3 data sources. Let's filter by confidence:
var targetWindow = aiSelectedDataSources.DetermineTargetWindow(TargetWindowStrategy.A_FEW_GOOD_ONES);
var threshold = aiSelectedDataSources.GetConfidenceThreshold(targetWindow);
//
// Filter the data sources by the threshold:
@@ -0,0 +1,96 @@
using System.Text;
using AIStudio.Chat;
namespace AIStudio.Tools.RAG;
public static class IRetrievalContextExtensions
{
private static readonly ILogger<IRetrievalContext> LOGGER = Program.SERVICE_PROVIDER.GetService<ILogger<IRetrievalContext>>()!;
public static async Task<string> AsMarkdown(this IReadOnlyList<IRetrievalContext> retrievalContexts, StringBuilder? sb = null, CancellationToken token = default)
{
sb ??= new StringBuilder();
var index = 0;
foreach(var retrievalContext in retrievalContexts)
{
index++;
await retrievalContext.AsMarkdown(sb, index, retrievalContexts.Count, token);
}
return sb.ToString();
}
public static async Task<string> AsMarkdown(this IRetrievalContext retrievalContext, StringBuilder? sb = null, int index = -1, int numTotalRetrievalContexts = -1, CancellationToken token = default)
{
sb ??= new StringBuilder();
switch (index)
{
case > 0 when numTotalRetrievalContexts is -1:
sb.AppendLine($"# Retrieval context {index}");
break;
case > 0 when numTotalRetrievalContexts > 0:
sb.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
break;
default:
sb.AppendLine("# Retrieval context");
break;
}
sb.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
sb.AppendLine($"Content category: {retrievalContext.Category}");
sb.AppendLine($"Content type: {retrievalContext.Type}");
sb.AppendLine($"Content path: {retrievalContext.Path}");
if(retrievalContext.Links.Count > 0)
{
sb.AppendLine("Additional links:");
foreach(var link in retrievalContext.Links)
sb.AppendLine($"- {link}");
}
switch(retrievalContext)
{
case RetrievalTextContext textContext:
sb.AppendLine();
sb.AppendLine("Matched text content:");
sb.AppendLine("````");
sb.AppendLine(textContext.MatchedText);
sb.AppendLine("````");
if(textContext.SurroundingContent.Count > 0)
{
sb.AppendLine();
sb.AppendLine("Surrounding text content:");
foreach(var surrounding in textContext.SurroundingContent)
{
sb.AppendLine();
sb.AppendLine("````");
sb.AppendLine(surrounding);
sb.AppendLine("````");
}
}
break;
case RetrievalImageContext imageContext:
sb.AppendLine();
sb.AppendLine("Matched image content as base64-encoded data:");
sb.AppendLine("````");
sb.AppendLine(await imageContext.AsBase64(token));
sb.AppendLine("````");
break;
default:
LOGGER.LogWarning($"The retrieval content type '{retrievalContext.Type}' of data source '{retrievalContext.DataSourceName}' at location '{retrievalContext.Path}' is not supported yet.");
break;
}
sb.AppendLine();
return sb.ToString();
}
}
@@ -9,7 +9,7 @@ public static class RetrievalContentCategoryExtensions
/// </summary>
/// <param name="contentType">The content type yielded by the ERI server.</param>
/// <returns>The corresponding retrieval content category.</returns>
public static RetrievalContentCategory ToRetrievalContentCategory(ContentType contentType) => contentType switch
public static RetrievalContentCategory ToRetrievalContentCategory(this ContentType contentType) => contentType switch
{
ContentType.NONE => RetrievalContentCategory.NONE,
ContentType.UNKNOWN => RetrievalContentCategory.UNKNOWN,
@@ -0,0 +1,38 @@
namespace AIStudio.Tools;
/// <summary>
/// Represents a target window for the number of items to match a threshold.
/// </summary>
/// <param name="NumMinItems">The minimum number of items to match the threshold. Should be at least one and less than targetWindowMin.</param>
/// <param name="TargetWindowMin">The minimum number of items in the target window. Should be at least 2 and more than numMinItems.</param>
/// <param name="TargetWindowMax">The maximum number of items in the target window.</param>
public readonly record struct TargetWindow(int NumMinItems, int TargetWindowMin, int TargetWindowMax, float MinThreshold)
{
/// <summary>
/// Determines if the target window is valid.
/// </summary>
/// <returns>True when the target window is valid; otherwise, false.</returns>
public bool IsValid()
{
if(this.NumMinItems < 1)
return false;
if(this.TargetWindowMin < this.NumMinItems)
return false;
if(this.TargetWindowMax < this.TargetWindowMin)
return false;
if(this.MinThreshold is < 0f or > 1f)
return false;
return true;
}
/// <summary>
/// Determines if the number of items is inside the target window.
/// </summary>
/// <param name="numItems">The number of items to check.</param>
/// <returns>True when the number of items is inside the target window; otherwise, false.</returns>
public bool InsideWindow(int numItems) => numItems >= this.TargetWindowMin && numItems <= this.TargetWindowMax;
}
@@ -0,0 +1,19 @@
namespace AIStudio.Tools;
public enum TargetWindowStrategy
{
/// <summary>
/// Means no target window strategy, which will effectively return all items.
/// </summary>
NONE,
/// <summary>
/// Searches for two up-to-three items but at least one.
/// </summary>
A_FEW_GOOD_ONES,
/// <summary>
/// Searches for the top 10% items that are better than guessing, i.e., with confidence greater than 0.5f.
/// </summary>
TOP10_BETTER_THAN_GUESSING,
}