shibatch commited on
Commit
7aab2f4
·
verified ·
1 Parent(s): 4af09c3

Upload README.md with huggingface_hub

Browse files
Files changed (1) hide show
  1. README.md +39 -12
README.md CHANGED
@@ -66,16 +66,43 @@ No custom Gemma 3 modeling code is used.
66
  import torch
67
  from transformers import Gemma3ForCausalLM, PreTrainedTokenizerFast
68
 
69
- model_dir = "hf"
70
- tokenizer = PreTrainedTokenizerFast.from_pretrained(model_dir)
71
- model = Gemma3ForCausalLM.from_pretrained(model_dir)
72
- model.eval()
73
-
74
- ids = [tokenizer.bos_token_id] + tokenizer.encode("Once upon", add_special_tokens=False)
75
- input_ids = torch.tensor([ids], dtype=torch.long)
76
-
77
- with torch.no_grad():
78
- out = model.generate(input_ids=input_ids, max_new_tokens=64, do_sample=False)
79
-
80
- print(tokenizer.decode(out[0], skip_special_tokens=True))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
81
  ```
 
66
  import torch
67
  from transformers import Gemma3ForCausalLM, PreTrainedTokenizerFast
68
 
69
+ def main():
70
+ repo_id = "shibatch/tinygemma3-2m"
71
+
72
+ print("Loading tokenizer...")
73
+ tokenizer = PreTrainedTokenizerFast.from_pretrained(repo_id, subfolder="hf")
74
+
75
+ print("Loading Gemma3 model weights...")
76
+ device = "cuda" if torch.cuda.is_available() else "cpu"
77
+
78
+ model = Gemma3ForCausalLM.from_pretrained(
79
+ repo_id,
80
+ subfolder="hf",
81
+ torch_dtype=torch.bfloat16 if device == "cuda" else torch.float32,
82
+ ).to(device)
83
+ model.eval()
84
+
85
+ prompt = "Once upon"
86
+ print(f"\nInput prompt: {prompt}")
87
+
88
+ input_ids = tokenizer.encode(prompt, add_special_tokens=False)
89
+ input_ids = [tokenizer.bos_token_id] + input_ids
90
+ input_ids = torch.tensor([input_ids], dtype=torch.long, device=device)
91
+
92
+ with torch.no_grad():
93
+ outputs = model.generate(
94
+ input_ids,
95
+ max_new_tokens=100,
96
+ do_sample=False,
97
+ repetition_penalty=1.0,
98
+ top_p=1.0,
99
+ pad_token_id=tokenizer.pad_token_id or tokenizer.bos_token_id,
100
+ eos_token_id=tokenizer.eos_token_id,
101
+ )
102
+
103
+ generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)
104
+ print(f"Generated output: {generated_text}")
105
+
106
+ if __name__ == "__main__":
107
+ main()
108
  ```