rank.rsannotatedrank.rssource73 lines · 2.9 KB · raw
1use log::{debug, trace, warn};
2use jev_protocol::{Choice, ChoiceAnswer, Json, Key, ProtocolError, Questions, Response};

Ordering memories by how well they fit a message, as one Jev question. A Choice answer carries a probability for every option, so a whole shortlist is ordered by one question (Jev is limited to thirty questions a minute across the machine; a question per candidate would spend the minute on one message).

8pub struct Rerank {
9    pub questions: Questions,
10    key: Key<ChoiceAnswer>,
11    n: usize,
12}
14impl Rerank {

Needs two to 255 candidates; callers rank only when there is something to choose between.

16    pub fn new(candidates: &[String]) -> Result<Self, ProtocolError> {
17        debug!("building a rerank question over {} candidates", candidates.len());
18        let mut questions = Questions::new();
19        let key = questions.choice(
20            "best",
21            Choice::new(
22                Json::text(
23                    "The state holds a message a young child just said to a toy cat. Which of these things \
24                     the cat remembers would help it answer the child best?",
25                ),
26                candidates.iter().enumerate().map(|(i, c)| ((i + 1).to_string(), Some(Json::text(c)))),
27            )?,
28        )?;
29        Ok(Self { questions, key, n: candidates.len() })
30    }
32    pub fn state(query: &str) -> Json {
33        let json = serde_json::json!({ "message": query });
34        trace!("rerank state: query of {} chars", query.len());
35        Json::canonical(&json.to_string()).unwrap_or_else(|e| {
36            warn!("rerank state is not canonical JSON ({e:?}); sending it as plain text");
37            Json::text(query)
38        })
39    }

A probability per candidate, in the order they were given.

42    pub fn probabilities(&self, response: &Response) -> Vec<f32> {
43        let answer = response.get(self.key);
44        let mut out = vec![0.0f32; self.n];
45        for (label, p) in &answer.probabilities {
46            if let Some(i) = label.parse::<usize>().ok().and_then(|n| n.checked_sub(1)) {
47                if i < out.len() {
48                    out[i] = *p as f32;
49                } else {
50                    warn!("rerank: Jev returned a probability for candidate {} of {}", i + 1, out.len());
51                }
52            } else {
53                warn!("rerank: Jev returned an option label that is not a candidate number");
54            }
55        }
56        trace!("rerank: {} probabilities read", out.len());
57        out
58    }
59}
61#[cfg(test)]
62mod tests {
63    use super::*;
64
65    #[test]
66    fn two_to_many_candidates_build_and_fewer_do_not() {
67        let c = |n: usize| (0..n).map(|i| format!("fact {i}")).collect::<Vec<_>>();
68        assert!(Rerank::new(&c(2)).is_ok());
69        assert!(Rerank::new(&c(8)).is_ok());
70        assert!(Rerank::new(&c(1)).is_err());
71        assert!(Rerank::new(&[]).is_err());
72    }
73}