Filter

Filter functions and classes

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.


source

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
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.removed

source

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


source

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)