Views
No views yet
output/contrastive_gemma_frozen_512/checkpoint_step_64000google/embeddinggemma-300m1"""
2Example: Using Fine-tuned Gemma Model for Tag Classification
3
4This example shows how to use the fine-tuned Gemma model for
5persona-conditioned tag classification using our module abstractions.
6
7Installation:
8 pip install git+https://github.com/Pieces/TAG-module.git@main
9 # Or: pip install -e .
10
11To use this model:
12 1. Download model.pt and config.json from the Hub
13 2. Place them in a directory (e.g., ./checkpoint/)
14 3. Update checkpoint_path below
15"""
16
17import torch
18from pathlib import Path
19from playground.validate_from_checkpoint import (
20 load_trained_model,
21 encode_query,
22 encode_general_tags,
23 compute_ranked_tags,
24)
25
26# Load the fine-tuned model
27print("Loading fine-tuned Gemma model...")
28
29# Option 1: Load from local checkpoint (if you have it)
30checkpoint_path = "output/contrastive_gemma_frozen_512/checkpoint_step_64000"
31
32# Option 2: Load from downloaded Hub files
33# 1. Download model.pt and config.json from:
34# https://huggingface.co/Pieces/gemma-frozen-512-step64000
35# 2. Place in a directory and update path:
36# checkpoint_path = "./downloaded_checkpoint"
37
38if not Path(checkpoint_path).exists():
39 print(f"⚠ Checkpoint not found at: {checkpoint_path}")
40 print("Please download model.pt and config.json from the Hub and update checkpoint_path")
41 exit(1)
42
43model, config = load_trained_model(checkpoint_path, device="cpu")
44model.eval()
45print("✓ Model loaded!")
46
47# Example query with persona
48query_text = "How to implement OAuth2 authentication in a Python Flask API?"
49persona_text = "I'm a backend developer working on web APIs and microservices"
50
51# Candidate tags to rank
52candidate_tags = [
53 "python", "flask", "oauth2", "authentication", "api",
54 "security", "web-development", "jwt", "rest-api", "backend",
55 "microservices", "fastapi", "django"
56]
57
58print(f"\nQuery: {query_text}")
59print(f"Persona: {persona_text}")
60print(f"Candidate tags: {candidate_tags}\n")
61
62# Encode query using base branch (no persona conditioning)
63print("Encoding query (base branch)...")
64with torch.inference_mode():
65 query_emb_base = encode_query(
66 model=model,
67 query_text=query_text,
68 persona_text=persona_text,
69 connectivity="high",
70 branch="base",
71 max_length=512,
72 use_pretrained_backbone=False,
73 extraction_mode="full",
74 )
75
76# Encode query using personalized branch (with persona conditioning)
77print("Encoding query (personalized branch)...")
78with torch.inference_mode():
79 query_emb_personalized = encode_query(
80 model=model,
81 query_text=query_text,
82 persona_text=persona_text,
83 connectivity="high",
84 branch="personalized", # Use personalized branch
85 max_length=512,
86 use_pretrained_backbone=False,
87 extraction_mode="full",
88 )
89
90# Encode tags
91print("Encoding tags...")
92with torch.inference_mode():
93 tag_embs_base = encode_general_tags(
94 model=model,
95 general_tags=candidate_tags,
96 connectivity="high",
97 branch="base",
98 persona_text=persona_text,
99 max_length=512,
100 use_pretrained_backbone=False,
101 extraction_mode="full",
102 )
103
104 tag_embs_personalized = encode_general_tags(
105 model=model,
106 general_tags=candidate_tags,
107 connectivity="high",
108 branch="personalized", # Use personalized branch
109 persona_text=persona_text,
110 max_length=512,
111 use_pretrained_backbone=False,
112 extraction_mode="full",
113 )
114
115print(f"Query embedding shape: {query_emb_base.shape}")
116print(f"Tag embeddings shape: {tag_embs_base.shape}")
117
118# Rank tags using base branch
119print("\n" + "="*60)
120print("Rankings (Base Branch - No Persona Conditioning):")
121print("="*60)
122ranked_tags_base = compute_ranked_tags(
123 query_emb=query_emb_base,
124 pos_embs=torch.empty(0, model.config.backbone_embedding_dim),
125 neg_embs=torch.empty(0, model.config.backbone_embedding_dim),
126 general_embs=tag_embs_base,
127 positive_tags=[],
128 negative_tags=[],
129 general_tags=candidate_tags,
130)
131
132for tag, rank, label, score in ranked_tags_base[:5]:
133 print(f"{rank:2d}. {tag:20s} (score: {score:.4f})")
134
135# Rank tags using personalized branch
136print("\n" + "="*60)
137print("Rankings (Personalized Branch - With Persona Conditioning):")
138print("="*60)
139ranked_tags_personalized = compute_ranked_tags(
140 query_emb=query_emb_personalized,
141 pos_embs=torch.empty(0, model.config.backbone_embedding_dim),
142 neg_embs=torch.empty(0, model.config.backbone_embedding_dim),
143 general_embs=tag_embs_personalized,
144 positive_tags=[],
145 negative_tags=[],
146 general_tags=candidate_tags,
147)
148
149for tag, rank, label, score in ranked_tags_personalized[:5]:
150 print(f"{rank:2d}. {tag:20s} (score: {score:.4f})")
151
152print("\n" + "="*60)
153print("Example complete!")
154print("\nNote: Personalized branch adapts tag rankings based on the persona context.")
155
1561# Install the repository first
2pip install git+https://github.com/Pieces/TAG-module.git@main
3# Or for local development:
4pip install -e .
5
6# Run the example
7python gemma_finetuned_example.py1@software{{tag_module,
2 title = {{TAG Module: Persona-Conditioned Contrastive Learning for Tag Classification}},
3 author = {{Your Name}},
4 year = {{2025}},
5 url = {{https://github.com/yourusername/tag-module}}
6}}