BananaMind-2.1-MiniCoder / curriculum_coder_30b.py
Banaxi-Tech's picture
Add MiniCoder architecture and streamed trainer
c1c66b7 verified
Raw
History Blame Contribute Delete
2 kB
"""Exact fixed 30B-token data mixture for NanoCoder and MiniCoder."""
from __future__ import annotations
TOTAL_TOKENS = 30_000_000_000
SOURCE_KEYS = ("stack_v3", "fineweb_edu")
TARGET_SHARES = {"stack_v3": 0.75, "fineweb_edu": 0.25}
TARGET_SOURCE_TOKENS = {
"stack_v3": 22_500_000_000,
"fineweb_edu": 7_500_000_000,
}
SOURCES = {
"stack_v3": {
"label": "The Stack v3 train",
"dataset_id": "HuggingFaceCode/stack-v3-train",
"config_name": None,
"kind": "repository",
},
"fineweb_edu": {
"label": "FineWeb-Edu",
"dataset_id": "HuggingFaceFW/fineweb-edu",
"config_name": "sample-100BT",
"kind": "text",
},
}
class TokenCreditScheduler:
"""Deterministic weighted-fair scheduling in supervised tokens."""
def __init__(self):
self.credits = {key: 0.0 for key in SOURCE_KEYS}
self.consumed = {key: 0 for key in SOURCE_KEYS}
def choose(self, tokens_this_step: int) -> str:
for key in SOURCE_KEYS:
self.credits[key] += TARGET_SHARES[key] * tokens_this_step
selected = max(SOURCE_KEYS, key=lambda key: self.credits[key])
self.credits[selected] -= tokens_this_step
self.consumed[selected] += tokens_this_step
return selected
def state_dict(self) -> dict[str, dict[str, float | int]]:
return {
"credits": dict(self.credits),
"consumed": dict(self.consumed),
}
def load_state_dict(self, state: dict) -> None:
if set(state["credits"]) != set(SOURCE_KEYS):
raise ValueError("Checkpoint curriculum keys do not match")
self.credits = {
key: float(state["credits"][key]) for key in SOURCE_KEYS
}
self.consumed = {
key: int(state["consumed"][key]) for key in SOURCE_KEYS
}
__all__ = [
"SOURCES",
"SOURCE_KEYS",
"TARGET_SHARES",
"TARGET_SOURCE_TOKENS",
"TOTAL_TOKENS",
"TokenCreditScheduler",
]