def build_batch(cutoff=10.5):
d_emb = 128
n_emb = 100
np.random.seed(42)
embeddings = np.random.randn(n_emb+1, d_emb)
query = Query.from_minimal(embedding=embeddings[-1])
results = [Item.from_minimal(id=i, embedding=embeddings[i]) for i in range(n_emb)]
query.add_query_results(results)
batch = Batch(queries=[query])
expected_failures = [i.id for i in results if np.linalg.norm(i.embedding)>=cutoff]
return batch, expected_failures
class NormFilter():
def __init__(self, cutoff=10.5):
self.cutoff = cutoff
def __call__(self, inputs: List[Item]) -> List[FilterResponse]:
embeddings = np.array([i.embedding for i in inputs])
norms = np.linalg.norm(embeddings, axis=-1)
results = [FilterResponse(valid=i<self.cutoff, data={'norm':i}) for i in norms]
return results
filter_func = NormFilter()
filter_module = FilterModule(filter_func)
batch, fails = build_batch()
batch2 = filter_module(batch)
assert len(batch2.flatten_query_results(skip_removed=True)[1]) == len(batch[0])-len(fails)
for i in range(len(batch[0])):
result = batch[0][i]
if i in fails:
assert result.internal.removed
else:
assert not result.internal.removed
assert result.internal.parent_id == batch[0].id
batch, fails = build_batch()
filter_func = NormFilter(cutoff=-1)
filter_module = FilterModule(filter_func)
batch, fails = build_batch(cutoff=-1)
batch2 = filter_module(batch)
assert batch2[0].internal.removedFilter
The Filter step optionally removes Item results based on some boolean criteria. This can be used to remove undesirable results prior to scoring and updating. The filter step is formalized by the FilterFunction schema, which maps inputs List[Item] to outputs List[FilterResponse]
The FilterModule manages execution of a FilterFunction. The FilterModule gathers valid items, sends them to the FilterFunction, and processes the results.
FilterModule
FilterModule (function:Callable[[List[emb_opt.schemas.Item]],List[emb_opt .schemas.FilterResponse]])
Module - module base class
Given an input Batch, the Module: 1. gathers inputs to the function 2. executes the function 3. validates the results of the function with output_schema 4. scatters results back into the Batch
| Type | Details | |
|---|---|---|
| function | typing.Callable[[typing.List[emb_opt.schemas.Item]], typing.List[emb_opt.schemas.FilterResponse]] | filter function |
FilterPlugin
FilterPlugin ()
FilterPlugin - documentation for plugin functions to FilterFunction
A valid FilterFunction is any function that maps List[Item] to List[FilterResponse]. The inputs will be given as Item objects. The outputs can be either a list of FilterResponse objects or a list of valid json dictionaries that match the FilterResponse schema.
Item schema:
{ 'id' : Optional[Union[str, int]] 'item' : Optional[Any], 'embedding' : List[float], 'score' : None, # will be None at this stage 'data' : Optional[Dict], }
Input schema:
List[Item]
FilterResponse schema:
{ 'valid' : bool, 'data' : Optional[Dict], }
Output schema:
List[FilterResponse]
The CompositeFilterPlugin can be used to chain together a list of valid FilterFunction
CompositeFilterPlugin
CompositeFilterPlugin (functions:List[Callable[[List[emb_opt.schemas.Item ]],List[emb_opt.schemas.FilterResponse]]])
Initialize self. See help(type(self)) for accurate signature.
| Type | Details | |
|---|---|---|
| functions | typing.List[typing.Callable[[typing.List[emb_opt.schemas.Item]], typing.List[emb_opt.schemas.FilterResponse]]] | list of filter functions |
def build_batch(cutoff=10.5):
d_emb = 128
n_emb = 100
np.random.seed(42)
embeddings = np.random.randn(n_emb+1, d_emb)
query = Query.from_minimal(embedding=embeddings[-1])
results = [Item.from_minimal(id=i, embedding=embeddings[i]) for i in range(n_emb)]
query.add_query_results(results)
batch = Batch(queries=[query])
return batch
def norm_filter(inputs: List[Item], cutoff: float=10.5) -> List[FilterResponse]:
embeddings = np.array([i.embedding for i in inputs])
norms = np.linalg.norm(embeddings, axis=-1)
results = [FilterResponse(valid=i<cutoff, data={'norm':i}) for i in norms]
return results
def sum_filter(inputs: List[Item], cutoff: float=0.0) -> List[FilterResponse]:
embeddings = np.array([i.embedding for i in inputs])
sums = embeddings.sum(-1)
results = [FilterResponse(valid=i>cutoff, data={'sum':i}) for i in sums]
return results
filter_func = CompositeFilterPlugin([norm_filter, sum_filter])
filter_module = FilterModule(filter_func)
batch = build_batch()
batch2 = filter_module(batch)