Git-Repository für bestehenden archivdms-Code initialisiert, Branch-/Commit-Konvention (feature/<ticket>-<slug>-Branches, Ticket-Prefix in Commit-Nachricht) etabliert.
618 lines
21 KiB
Go
618 lines
21 KiB
Go
// Package classifier implements a dependency-free multinomial Naive-Bayes text
|
||
// classifier that supplements archivdms's rule-based matching engine
|
||
// (internal/matching) WITHOUT introducing any LLM dependency. It is designed to
|
||
// give clean, self-standing results even when no Ollama/LLM is configured: the
|
||
// tokenizer, German stop-word handling and Laplace smoothing are the quality
|
||
// levers and are treated as first-class, not minimal.
|
||
//
|
||
// The model is persisted in the ml_classifier_tokens / ml_classifier_classes
|
||
// tables (see internal/storage/ml_classifier.go). Training is a full
|
||
// per-tenant/per-kind rebuild (DELETE + bulk insert) — deliberately simple, no
|
||
// incremental updates. Classification (Predict) reads that persisted model and
|
||
// returns softmax-normalised posterior probabilities so the caller can apply a
|
||
// single confidence floor comparable to the heuristic provider's
|
||
// suggestionFloor.
|
||
//
|
||
// This package MUST NOT import internal/storage (storage imports it, for
|
||
// Store.GenerateNaiveBayesSuggestions) — it talks to Postgres through the small
|
||
// DB interface below, which *pgxpool.Pool satisfies.
|
||
package classifier
|
||
|
||
import (
|
||
"context"
|
||
"fmt"
|
||
"math"
|
||
"sort"
|
||
"strings"
|
||
"unicode"
|
||
|
||
"github.com/jackc/pgx/v5"
|
||
"github.com/jackc/pgx/v5/pgconn"
|
||
)
|
||
|
||
// DB is the minimal Postgres surface the classifier needs. *pgxpool.Pool (and
|
||
// pgx.Tx, for the value passed to Predict from within a transaction) satisfy it.
|
||
type DB interface {
|
||
Begin(ctx context.Context) (pgx.Tx, error)
|
||
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
|
||
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
||
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
|
||
}
|
||
|
||
const (
|
||
// MinDocsPerClass is the minimum number of labelled training documents a
|
||
// class (document_type / correspondent / tag) must have before it is kept in
|
||
// the model. Below this the class is silently skipped (no error): too few
|
||
// examples produce an over-confident, unreliable token distribution.
|
||
MinDocsPerClass = 20
|
||
|
||
// laplaceAlpha is the additive (Laplace/Lidstone) smoothing constant applied
|
||
// to every token count. alpha=1 (classic add-one) keeps unseen tokens from
|
||
// zeroing out a whole class's likelihood while staying conservative for
|
||
// small vocabularies.
|
||
laplaceAlpha = 1.0
|
||
|
||
// SuggestionFloor is the minimum softmax posterior probability at which a
|
||
// class is surfaced as a candidate. Chosen to mirror the heuristic
|
||
// provider's suggestionFloor (0.55) so the two engines feel consistent to
|
||
// the user: below this the model is essentially undecided between classes.
|
||
SuggestionFloor = 0.55
|
||
|
||
// maxCandidates caps how many classes Predict returns, sorted by posterior
|
||
// descending (matches maxSuggestionCandidates in the storage layer).
|
||
maxCandidates = 5
|
||
|
||
// maxExplanationTokens is how many "most decisive" tokens are attached to a
|
||
// candidate as a human-readable explanation.
|
||
maxExplanationTokens = 5
|
||
|
||
// minTokenRunes / maxTokenRunes bound token length: 1-rune tokens are noise;
|
||
// absurdly long runs are almost always OCR garbage (barcodes, scan
|
||
// artefacts) rather than meaningful words.
|
||
minTokenRunes = 2
|
||
maxTokenRunes = 40
|
||
|
||
// minNumericTokenRunes: purely numeric tokens shorter than this (page
|
||
// numbers, single amounts, "1."/"12") are dropped as noise, but longer
|
||
// numeric runs are KEPT — invoice/customer numbers, IBAN fragments and years
|
||
// recur across a correspondent's documents and carry real signal.
|
||
minNumericTokenRunes = 4
|
||
)
|
||
|
||
// Kinds are the three classifiable taxonomy kinds. Mirrors the CHECK constraint
|
||
// on ml_classifier_tokens.kind.
|
||
const (
|
||
KindDocumentTypes = "document_types"
|
||
KindCorrespondents = "correspondents"
|
||
KindTags = "tags"
|
||
)
|
||
|
||
// SuggestionCandidate is one scored class prediction. Score is a softmax
|
||
// posterior probability in [0,1]. TopTokens lists (at most maxExplanationTokens)
|
||
// input tokens that contributed most to selecting this class over the runner-up
|
||
// — the model's explanation for GoBD-Nachvollziehbarkeit / user trust.
|
||
type SuggestionCandidate struct {
|
||
EntityID int64 `json:"entity_id"`
|
||
Score float64 `json:"score"`
|
||
TopTokens []string `json:"top_tokens"`
|
||
}
|
||
|
||
// Classifier is a thin, stateless wrapper around a DB handle. Safe to construct
|
||
// per call.
|
||
type Classifier struct {
|
||
db DB
|
||
}
|
||
|
||
// New returns a Classifier backed by db.
|
||
func New(db DB) *Classifier {
|
||
return &Classifier{db: db}
|
||
}
|
||
|
||
// validKind guards the kind against the allowlist so it can be interpolated
|
||
// nowhere (all queries parameterise it) but callers still fail fast on typos.
|
||
func validKind(kind string) error {
|
||
switch kind {
|
||
case KindDocumentTypes, KindCorrespondents, KindTags:
|
||
return nil
|
||
default:
|
||
return fmt.Errorf("classifier: unknown kind %q", kind)
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// Tokenizer
|
||
// ---------------------------------------------------------------------------
|
||
|
||
// germanStopwords is a curated list of high-frequency German function words
|
||
// (plus a few ubiquitous document-boilerplate terms) that carry no class
|
||
// signal. Removing them sharpens the per-class token distributions. Kept as a
|
||
// set for O(1) lookup. Not exhaustive by design — the goal is to strip the
|
||
// worst offenders, not to stem the language.
|
||
var germanStopwords = func() map[string]struct{} {
|
||
words := []string{
|
||
"der", "die", "das", "den", "dem", "des", "ein", "eine", "einen", "einem",
|
||
"einer", "eines", "und", "oder", "aber", "auch", "sich", "nicht", "mit",
|
||
"für", "von", "vom", "im", "in", "am", "an", "auf", "aus", "bei", "bis",
|
||
"durch", "gegen", "ohne", "um", "unter", "über", "zwischen", "nach", "zu",
|
||
"zur", "zum", "vor", "hinter", "neben", "ist", "sind", "war", "waren",
|
||
"wird", "werden", "wurde", "wurden", "sein", "seine", "seiner", "ihre",
|
||
"ihrer", "ihren", "haben", "hat", "hatte", "hatten", "kann", "können",
|
||
"muss", "müssen", "soll", "sollen", "als", "wie", "wenn", "dann", "dass",
|
||
"daß", "weil", "denn", "doch", "nur", "noch", "schon", "sehr", "hier",
|
||
"dort", "man", "wir", "sie", "ich", "du", "er", "es", "ihr", "uns",
|
||
"euch", "mein", "dein", "unser", "diese", "dieser", "dieses", "diesem",
|
||
"jede", "jeder", "jedes", "alle", "allen", "aller", "kein", "keine",
|
||
"keinen", "mehr", "sehr", "so", "auch", "wieder", "bitte", "danke",
|
||
"gmbh", "seite", "www", "http", "https", "email", "mail", "tel",
|
||
}
|
||
m := make(map[string]struct{}, len(words))
|
||
for _, w := range words {
|
||
m[w] = struct{}{}
|
||
}
|
||
return m
|
||
}()
|
||
|
||
// tokenize normalises text into a bag of meaningful tokens:
|
||
// - Unicode-aware lowercasing.
|
||
// - A token is a maximal run of letters and/or digits (so "de89370400"
|
||
// survives as one token, and dots/slashes/whitespace all split). This keeps
|
||
// structured identifiers (IBAN fragments, invoice numbers) intact instead of
|
||
// shredding them into single digits.
|
||
// - German stop-words are dropped.
|
||
// - Tokens shorter than minTokenRunes or longer than maxTokenRunes are dropped.
|
||
// - Purely numeric tokens shorter than minNumericTokenRunes are dropped
|
||
// (page numbers, trivial amounts), longer ones are kept (they recur and
|
||
// carry signal). Mixed letter+digit tokens are always kept.
|
||
//
|
||
// Returns a frequency map (token -> count in this text), which is exactly what
|
||
// the multinomial model consumes.
|
||
func tokenize(text string) map[string]int {
|
||
freq := make(map[string]int)
|
||
var b strings.Builder
|
||
flush := func() {
|
||
if b.Len() == 0 {
|
||
return
|
||
}
|
||
tok := b.String()
|
||
b.Reset()
|
||
addToken(freq, tok)
|
||
}
|
||
for _, r := range text {
|
||
if unicode.IsLetter(r) || unicode.IsDigit(r) {
|
||
b.WriteRune(unicode.ToLower(r))
|
||
continue
|
||
}
|
||
flush()
|
||
}
|
||
flush()
|
||
return freq
|
||
}
|
||
|
||
// addToken applies the length / numeric / stop-word filters and, if the token
|
||
// survives, increments its frequency.
|
||
func addToken(freq map[string]int, tok string) {
|
||
runes := []rune(tok)
|
||
if len(runes) < minTokenRunes || len(runes) > maxTokenRunes {
|
||
return
|
||
}
|
||
if _, stop := germanStopwords[tok]; stop {
|
||
return
|
||
}
|
||
if isAllDigits(runes) && len(runes) < minNumericTokenRunes {
|
||
return
|
||
}
|
||
freq[tok]++
|
||
}
|
||
|
||
func isAllDigits(runes []rune) bool {
|
||
for _, r := range runes {
|
||
if !unicode.IsDigit(r) {
|
||
return false
|
||
}
|
||
}
|
||
return true
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// Training
|
||
// ---------------------------------------------------------------------------
|
||
|
||
// classAccum accumulates token statistics for one class during training.
|
||
type classAccum struct {
|
||
docCount int64
|
||
totalTokens int64
|
||
tokens map[string]int64
|
||
}
|
||
|
||
// Train rebuilds the Naive-Bayes model for one tenant and one kind from scratch.
|
||
// It reads every labelled, non-deleted training document (assignments made
|
||
// manually OR by the rule engine — NOT prior ml_accepted ones, to avoid the
|
||
// model reinforcing its own past guesses), tokenizes the title+OCR text, counts
|
||
// tokens per class, drops classes with fewer than MinDocsPerClass documents, and
|
||
// writes the result with a DELETE + bulk COPY inside a single transaction
|
||
// (previous model for this tenant/kind is fully replaced).
|
||
//
|
||
// Returns the number of training documents actually used (across retained
|
||
// classes). A kind with no qualifying data trains to an empty model and returns
|
||
// 0 — this is not an error.
|
||
func (c *Classifier) Train(ctx context.Context, tenantID int64, kind string) (docCount int, err error) {
|
||
if err := validKind(kind); err != nil {
|
||
return 0, err
|
||
}
|
||
|
||
rows, err := c.db.Query(ctx, trainingQuery(kind), tenantID)
|
||
if err != nil {
|
||
return 0, fmt.Errorf("classifier: read training data (%s): %w", kind, err)
|
||
}
|
||
defer rows.Close()
|
||
|
||
classes := make(map[int64]*classAccum)
|
||
for rows.Next() {
|
||
var entityID int64
|
||
var title, ocr string
|
||
if err := rows.Scan(&entityID, &title, &ocr); err != nil {
|
||
return 0, fmt.Errorf("classifier: scan training row (%s): %w", kind, err)
|
||
}
|
||
acc := classes[entityID]
|
||
if acc == nil {
|
||
acc = &classAccum{tokens: make(map[string]int64)}
|
||
classes[entityID] = acc
|
||
}
|
||
acc.docCount++
|
||
for tok, n := range tokenize(title + "\n" + ocr) {
|
||
acc.tokens[tok] += int64(n)
|
||
acc.totalTokens += int64(n)
|
||
}
|
||
}
|
||
if err := rows.Err(); err != nil {
|
||
return 0, fmt.Errorf("classifier: iterate training rows (%s): %w", kind, err)
|
||
}
|
||
|
||
// Keep only classes with enough evidence.
|
||
retained := make(map[int64]*classAccum)
|
||
used := 0
|
||
for id, acc := range classes {
|
||
if acc.docCount < MinDocsPerClass {
|
||
continue
|
||
}
|
||
retained[id] = acc
|
||
used += int(acc.docCount)
|
||
}
|
||
|
||
if err := c.persist(ctx, tenantID, kind, retained); err != nil {
|
||
return 0, err
|
||
}
|
||
return used, nil
|
||
}
|
||
|
||
// trainingQuery returns the SQL that yields (entity_id, title, ocr_text) rows
|
||
// for a kind, restricted to manual/rule-assigned, non-deleted documents.
|
||
func trainingQuery(kind string) string {
|
||
switch kind {
|
||
case KindDocumentTypes:
|
||
return `
|
||
SELECT d.doc_type_id, d.title, COALESCE(d.ocr_text, '')
|
||
FROM documents d
|
||
WHERE d.tenant_id = $1
|
||
AND d.deleted_at IS NULL
|
||
AND d.doc_type_id IS NOT NULL
|
||
AND d.doc_type_assigned_via IN ('manual','rule')`
|
||
case KindCorrespondents:
|
||
return `
|
||
SELECT d.correspondent_id, d.title, COALESCE(d.ocr_text, '')
|
||
FROM documents d
|
||
WHERE d.tenant_id = $1
|
||
AND d.deleted_at IS NULL
|
||
AND d.correspondent_id IS NOT NULL
|
||
AND d.correspondent_assigned_via IN ('manual','rule')`
|
||
case KindTags:
|
||
return `
|
||
SELECT dt.tag_id, d.title, COALESCE(d.ocr_text, '')
|
||
FROM document_tags dt
|
||
JOIN documents d ON d.id = dt.document_id
|
||
WHERE d.tenant_id = $1
|
||
AND d.deleted_at IS NULL
|
||
AND dt.assigned_via IN ('manual','rule')`
|
||
default:
|
||
return ""
|
||
}
|
||
}
|
||
|
||
// persist replaces the stored model for tenant/kind with the retained classes,
|
||
// atomically (DELETE + COPY inside one transaction). An empty retained map still
|
||
// clears the previous model — a class that dropped below the threshold must not
|
||
// keep serving stale predictions.
|
||
func (c *Classifier) persist(ctx context.Context, tenantID int64, kind string, retained map[int64]*classAccum) error {
|
||
tx, err := c.db.Begin(ctx)
|
||
if err != nil {
|
||
return fmt.Errorf("classifier: begin tx: %w", err)
|
||
}
|
||
defer tx.Rollback(ctx)
|
||
|
||
if _, err := tx.Exec(ctx, `DELETE FROM ml_classifier_tokens WHERE tenant_id = $1 AND kind = $2`, tenantID, kind); err != nil {
|
||
return fmt.Errorf("classifier: clear tokens: %w", err)
|
||
}
|
||
if _, err := tx.Exec(ctx, `DELETE FROM ml_classifier_classes WHERE tenant_id = $1 AND kind = $2`, tenantID, kind); err != nil {
|
||
return fmt.Errorf("classifier: clear classes: %w", err)
|
||
}
|
||
|
||
classRows := make([][]any, 0, len(retained))
|
||
tokenRows := make([][]any, 0)
|
||
for entityID, acc := range retained {
|
||
classRows = append(classRows, []any{tenantID, kind, entityID, acc.docCount, acc.totalTokens})
|
||
for tok, cnt := range acc.tokens {
|
||
tokenRows = append(tokenRows, []any{tenantID, kind, entityID, tok, cnt})
|
||
}
|
||
}
|
||
|
||
if len(classRows) > 0 {
|
||
if _, err := tx.CopyFrom(ctx,
|
||
pgx.Identifier{"ml_classifier_classes"},
|
||
[]string{"tenant_id", "kind", "entity_id", "doc_count", "total_tokens"},
|
||
pgx.CopyFromRows(classRows)); err != nil {
|
||
return fmt.Errorf("classifier: copy classes: %w", err)
|
||
}
|
||
}
|
||
if len(tokenRows) > 0 {
|
||
if _, err := tx.CopyFrom(ctx,
|
||
pgx.Identifier{"ml_classifier_tokens"},
|
||
[]string{"tenant_id", "kind", "entity_id", "token", "count"},
|
||
pgx.CopyFromRows(tokenRows)); err != nil {
|
||
return fmt.Errorf("classifier: copy tokens: %w", err)
|
||
}
|
||
}
|
||
|
||
if err := tx.Commit(ctx); err != nil {
|
||
return fmt.Errorf("classifier: commit tx: %w", err)
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// Prediction
|
||
// ---------------------------------------------------------------------------
|
||
|
||
// classStat holds the persisted per-class statistics needed for scoring.
|
||
type classStat struct {
|
||
entityID int64
|
||
docCount int64
|
||
totalTokens int64
|
||
logLikeAt float64 // running log-likelihood accumulator (built during Predict)
|
||
}
|
||
|
||
// Predict scores the given text against the trained model for tenant/kind and
|
||
// returns the classes whose softmax posterior probability is at least
|
||
// SuggestionFloor, top maxCandidates, sorted by score descending. Each candidate
|
||
// carries its most decisive tokens as an explanation.
|
||
//
|
||
// Scoring is the standard multinomial Naive-Bayes log-likelihood with Laplace
|
||
// smoothing:
|
||
//
|
||
// logscore(c) = log P(c) + Σ_t freq(t) · log( (count(t,c)+α) / (Σtokens_c + α·V) )
|
||
//
|
||
// where V is the tenant/kind vocabulary size. The log-scores are then softmaxed
|
||
// (max-subtracted for numerical stability) into posterior probabilities so a
|
||
// single, interpretable confidence floor can be applied.
|
||
func (c *Classifier) Predict(ctx context.Context, tenantID int64, kind string, text string) ([]SuggestionCandidate, error) {
|
||
if err := validKind(kind); err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
stats, totalDocs, err := c.loadClasses(ctx, tenantID, kind)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if len(stats) == 0 || totalDocs == 0 {
|
||
return []SuggestionCandidate{}, nil // untrained kind: no suggestions, not an error
|
||
}
|
||
|
||
vocab, err := c.vocabSize(ctx, tenantID, kind)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
freq := tokenize(text)
|
||
if len(freq) == 0 {
|
||
return []SuggestionCandidate{}, nil
|
||
}
|
||
tokens := make([]string, 0, len(freq))
|
||
for t := range freq {
|
||
tokens = append(tokens, t)
|
||
}
|
||
|
||
// tokenCounts[token][entityID] = stored count.
|
||
tokenCounts, err := c.loadTokenCounts(ctx, tenantID, kind, tokens)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
// Per-token, per-class smoothed log-probability, plus per-class log score.
|
||
// perTokenLog[token][entityID] retained for the explanation step.
|
||
perTokenLog := make(map[string]map[int64]float64, len(tokens))
|
||
for i := range stats {
|
||
st := &stats[i]
|
||
st.logLikeAt = math.Log(float64(st.docCount) / float64(totalDocs)) // log prior
|
||
}
|
||
denom := make(map[int64]float64, len(stats))
|
||
for i := range stats {
|
||
st := &stats[i]
|
||
denom[st.entityID] = float64(st.totalTokens) + laplaceAlpha*float64(vocab)
|
||
}
|
||
|
||
for _, tok := range tokens {
|
||
perClass := tokenCounts[tok]
|
||
logs := make(map[int64]float64, len(stats))
|
||
for i := range stats {
|
||
st := &stats[i]
|
||
cnt := float64(perClass[st.entityID]) // 0 if unseen
|
||
logp := math.Log((cnt + laplaceAlpha) / denom[st.entityID])
|
||
logs[st.entityID] = logp
|
||
st.logLikeAt += float64(freq[tok]) * logp
|
||
}
|
||
perTokenLog[tok] = logs
|
||
}
|
||
|
||
// Softmax over the class log-scores.
|
||
scores := softmax(stats)
|
||
|
||
cands := make([]SuggestionCandidate, 0, len(stats))
|
||
// Rank runner-up for the explanation (second-highest posterior).
|
||
for i := range stats {
|
||
st := stats[i]
|
||
p := scores[st.entityID]
|
||
if p < SuggestionFloor {
|
||
continue
|
||
}
|
||
cands = append(cands, SuggestionCandidate{
|
||
EntityID: st.entityID,
|
||
Score: p,
|
||
TopTokens: decisiveTokens(st.entityID, stats, freq, perTokenLog),
|
||
})
|
||
}
|
||
sort.SliceStable(cands, func(i, j int) bool { return cands[i].Score > cands[j].Score })
|
||
if len(cands) > maxCandidates {
|
||
cands = cands[:maxCandidates]
|
||
}
|
||
return cands, nil
|
||
}
|
||
|
||
// softmax converts the per-class log scores into posterior probabilities,
|
||
// subtracting the max log score first for numerical stability.
|
||
func softmax(stats []classStat) map[int64]float64 {
|
||
maxLog := math.Inf(-1)
|
||
for i := range stats {
|
||
if stats[i].logLikeAt > maxLog {
|
||
maxLog = stats[i].logLikeAt
|
||
}
|
||
}
|
||
sum := 0.0
|
||
exp := make(map[int64]float64, len(stats))
|
||
for i := range stats {
|
||
e := math.Exp(stats[i].logLikeAt - maxLog)
|
||
exp[stats[i].entityID] = e
|
||
sum += e
|
||
}
|
||
out := make(map[int64]float64, len(stats))
|
||
if sum == 0 {
|
||
return out
|
||
}
|
||
for id, e := range exp {
|
||
out[id] = e / sum
|
||
}
|
||
return out
|
||
}
|
||
|
||
// decisiveTokens returns the (up to maxExplanationTokens) input tokens that most
|
||
// favoured winner over the strongest competing class, weighted by their
|
||
// frequency in the text. Positive margin = the token pushed toward winner.
|
||
func decisiveTokens(winner int64, stats []classStat, freq map[string]int, perTokenLog map[string]map[int64]float64) []string {
|
||
type scored struct {
|
||
token string
|
||
margin float64
|
||
}
|
||
out := make([]scored, 0, len(freq))
|
||
for tok, logs := range perTokenLog {
|
||
winLog, ok := logs[winner]
|
||
if !ok {
|
||
continue
|
||
}
|
||
// Best competing class's log-prob for this token.
|
||
competitor := math.Inf(-1)
|
||
for _, st := range stats {
|
||
if st.entityID == winner {
|
||
continue
|
||
}
|
||
if l := logs[st.entityID]; l > competitor {
|
||
competitor = l
|
||
}
|
||
}
|
||
if math.IsInf(competitor, -1) {
|
||
competitor = winLog // single-class case: no margin
|
||
}
|
||
margin := float64(freq[tok]) * (winLog - competitor)
|
||
if margin <= 0 {
|
||
continue
|
||
}
|
||
out = append(out, scored{token: tok, margin: margin})
|
||
}
|
||
sort.SliceStable(out, func(i, j int) bool { return out[i].margin > out[j].margin })
|
||
tokens := make([]string, 0, maxExplanationTokens)
|
||
for _, s := range out {
|
||
if len(tokens) >= maxExplanationTokens {
|
||
break
|
||
}
|
||
tokens = append(tokens, s.token)
|
||
}
|
||
return tokens
|
||
}
|
||
|
||
// loadClasses reads the persisted per-class stats for tenant/kind and the total
|
||
// document count across them (the denominator of the class priors).
|
||
func (c *Classifier) loadClasses(ctx context.Context, tenantID int64, kind string) ([]classStat, int64, error) {
|
||
rows, err := c.db.Query(ctx,
|
||
`SELECT entity_id, doc_count, total_tokens
|
||
FROM ml_classifier_classes
|
||
WHERE tenant_id = $1 AND kind = $2`, tenantID, kind)
|
||
if err != nil {
|
||
return nil, 0, fmt.Errorf("classifier: load classes: %w", err)
|
||
}
|
||
defer rows.Close()
|
||
|
||
var out []classStat
|
||
var totalDocs int64
|
||
for rows.Next() {
|
||
var st classStat
|
||
if err := rows.Scan(&st.entityID, &st.docCount, &st.totalTokens); err != nil {
|
||
return nil, 0, fmt.Errorf("classifier: scan class: %w", err)
|
||
}
|
||
totalDocs += st.docCount
|
||
out = append(out, st)
|
||
}
|
||
return out, totalDocs, rows.Err()
|
||
}
|
||
|
||
// vocabSize returns the number of distinct tokens in the tenant/kind model — the
|
||
// V in Laplace smoothing.
|
||
func (c *Classifier) vocabSize(ctx context.Context, tenantID int64, kind string) (int64, error) {
|
||
var v int64
|
||
if err := c.db.QueryRow(ctx,
|
||
`SELECT COUNT(DISTINCT token) FROM ml_classifier_tokens WHERE tenant_id = $1 AND kind = $2`,
|
||
tenantID, kind).Scan(&v); err != nil {
|
||
return 0, fmt.Errorf("classifier: vocab size: %w", err)
|
||
}
|
||
return v, nil
|
||
}
|
||
|
||
// loadTokenCounts fetches the per-class counts for exactly the input tokens
|
||
// (one query, token = ANY($3)), returning token -> entityID -> count.
|
||
func (c *Classifier) loadTokenCounts(ctx context.Context, tenantID int64, kind string, tokens []string) (map[string]map[int64]int64, error) {
|
||
out := make(map[string]map[int64]int64, len(tokens))
|
||
if len(tokens) == 0 {
|
||
return out, nil
|
||
}
|
||
rows, err := c.db.Query(ctx,
|
||
`SELECT token, entity_id, count
|
||
FROM ml_classifier_tokens
|
||
WHERE tenant_id = $1 AND kind = $2 AND token = ANY($3)`,
|
||
tenantID, kind, tokens)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("classifier: load token counts: %w", err)
|
||
}
|
||
defer rows.Close()
|
||
|
||
for rows.Next() {
|
||
var tok string
|
||
var entityID, cnt int64
|
||
if err := rows.Scan(&tok, &entityID, &cnt); err != nil {
|
||
return nil, fmt.Errorf("classifier: scan token count: %w", err)
|
||
}
|
||
m := out[tok]
|
||
if m == nil {
|
||
m = make(map[int64]int64)
|
||
out[tok] = m
|
||
}
|
||
m[entityID] = cnt
|
||
}
|
||
return out, rows.Err()
|
||
}
|