diff --git a/packages/atproto/classifiers/tfjs/index.ts b/packages/atproto/classifiers/tfjs/index.ts index 058cb45..b74c39f 100644 --- a/packages/atproto/classifiers/tfjs/index.ts +++ b/packages/atproto/classifiers/tfjs/index.ts @@ -5,17 +5,23 @@ import { type ExtractedText, } from "../../domain/extract-text-from-post"; import { postTexts } from "../../domain/post/post-texts.table"; -import type { FeedPostWithUri } from "../../domain/queue-post"; import type { ClassifierFn } from "../types"; import { classify } from "./classify"; import { isModelOutdated } from "./is-model-outdated"; import { train } from "./train"; +import type { FeedPostWithUri } from "@andrioid/jetstream"; +import { trainingSetSize } from "../../domain/training-set-size"; -type FnType = (ctx: AtContext) => Promise; +type FnType = (ctx: AtContext) => Promise; +const MINIMUM_POST_COUNT = 50; export const createBayesClassiferFn: FnType = async (ctx) => { // 1. Check if current model is too old - const isTooOld = await isModelOutdated(ctx); + const sizeOfSet = await trainingSetSize(ctx); + if (sizeOfSet < MINIMUM_POST_COUNT) { + console.log("[classifier] not enough training data, aborting tfjsbayes"); + return; + } //const loader = isTooOld ? train : loadModelFromDb; const loader = train; const m = await loader(ctx); diff --git a/packages/atproto/classifiers/types.ts b/packages/atproto/classifiers/types.ts index 71c4ea5..8fec9dc 100644 --- a/packages/atproto/classifiers/types.ts +++ b/packages/atproto/classifiers/types.ts @@ -1,5 +1,5 @@ +import type { FeedPostWithUri } from "@andrioid/jetstream"; import type { AtContext } from "../context"; -import type { FeedPostWithUri } from "../domain/queue-post"; export type ClassifierFn = (args: { ctx: AtContext; diff --git a/packages/atproto/domain/get-tech-all-feed.ts b/packages/atproto/domain/get-tech-all-feed.ts index 09d6b7a..b5d80b2 100644 --- a/packages/atproto/domain/get-tech-all-feed.ts +++ b/packages/atproto/domain/get-tech-all-feed.ts @@ -21,14 +21,11 @@ export async function getTechAllFeed( .leftJoin(followTable, eq(postTable.authorId, followTable.follows)) .innerJoin( postScores, - and( - eq(postScores.postId, postTable.id), - eq(postScores.tagId, "tech"), - cursor ? gt(postTable.created, fromCursor(cursor)) : undefined - ) + and(eq(postScores.postId, postTable.id), eq(postScores.tagId, "tech")) ) .where( and( + cursor ? gt(postTable.created, fromCursor(cursor)) : undefined, gt(postTable.flags, 0), gte(postScores.avgScore, 75), // TODO 80 or( @@ -38,6 +35,7 @@ export async function getTechAllFeed( ) ) .orderBy(desc(postTable.created)) + .groupBy(postTable.id) .limit(30); diff --git a/packages/atproto/domain/training-set-size.ts b/packages/atproto/domain/training-set-size.ts new file mode 100644 index 0000000..b5ca613 --- /dev/null +++ b/packages/atproto/domain/training-set-size.ts @@ -0,0 +1,16 @@ +import { and, eq } from "drizzle-orm"; +import type { AtContext } from "../context"; +import { postTags } from "./post/post-tag.table"; + +export async function trainingSetSize(ctx: AtContext) { + const allTrainingPosts = await ctx.db + .select({ + postId: postTags.postId, + tag: postTags.tagId, + score: postTags.score, + }) + .from(postTags) + .where(and(eq(postTags.algo, "manual"), eq(postTags.tagId, "tech"))); + + return allTrainingPosts.length; +} diff --git a/packages/atproto/scripts/classifier.ts b/packages/atproto/scripts/classifier.ts index c77acec..9cc5197 100644 --- a/packages/atproto/scripts/classifier.ts +++ b/packages/atproto/scripts/classifier.ts @@ -26,9 +26,11 @@ export async function classifier() { mutedWordsClassifier, techWordsRegExp, techLinkClassifier, - await createBayesClassiferFn(ctx), ]; + const bayes = await createBayesClassiferFn(ctx); + if (bayes) classifiers.push(bayes); + const res = await ctx.db .select() .from(postTable)