Hierarchical classification in one API call

A product catalog has dozens of categories, grouped into departments. Instead of one long list, classify in two steps: pick the department, then a category inside it.

With TypeLLM, both steps happen in one API call. The department's categories are only asked when that department is chosen, so each step chooses from a short list.

The single call is in the Classify into a category tree example. This recipe builds the questions from a tree of 64 categories, then classifies a list of products with them.

Set up

Install the library:

pip install "typellm>=0.6.7"

Get an API key from your dashboard:

from typellm import TypeLLMClient

client = TypeLLMClient(api_key="YOUR_API_KEY")

PRODUCTS is a list of 24 one-line product descriptions; see the notebook for details.

Write the category tree

Eight departments with eight categories each: 64 categories in all.

TREE = {
    "electronics": ["phones", "laptops", "audio", "cameras",
                    "tv_video", "wearables", "smart_home", "accessories"],
    "home_kitchen": ["cookware", "coffee_tea", "small_appliances", "bedding",
                     "furniture", "storage", "lighting", "decor"],
    # ... and clothing, sports_outdoors, beauty, toys_games, office, pet_supplies
}

Build the questions from the tree

department chooses from the departments. Each department gets its own field with its categories, asked only when the department matches. However large the tree grows, this code stays the same.

def build_questions(tree):
    questions = {"department": {
        "type": "string", "enum": list(tree),
        "instructions": "Which department does this product belong to?",
    }}
    for department, categories in tree.items():
        questions[department] = {
            "type": "string", "enum": categories, "when": {"department": department},
            "instructions": f"Which {department.replace('_', ' ')} category does this product belong to?",
        }
    return questions

QUESTIONS = build_questions(TREE)

Classify in one API call

The answer holds the department and the one category field that ran:

def classify(product):
    result = client.generate(context=f"Product: {product}", questions=QUESTIONS).result
    department = result["department"]
    return department, result[department]
electronics      audio            <- Noise-cancelling headphones
home_kitchen     coffee_tea       <- Stainless steel French press
clothing         shoes            <- Leather ankle boots
sports_outdoors  camping          <- Two-person backpacking tent
beauty           haircare         <- Sulfate-free shampoo
toys_games       remote_control   <- Remote control off-road truck
office           labels           <- Label maker
pet_supplies     aquarium         <- Aquarium filter

Final result

electronics      laptops, audio, cameras
home_kitchen     cookware, coffee_tea, bedding
clothing         bottoms, outerwear, shoes
sports_outdoors  camping, cycling, fitness
beauty           skincare, makeup, haircare
toys_games       puzzles, board_games, remote_control
office           writing, office_furniture, labels
pet_supplies     dog_food, litter, aquarium

All 24 products land in the department and category you would expect. The notebook runs all of it with your key: open it in Colab.

← all recipes