feat: GPU 생성 작업 샤딩 지원

This commit is contained in:
hbyang 2026-08-20 17:19:41 +09:00
parent 78e9f8648a
commit d0b5876ae5
2 changed files with 20 additions and 1 deletions

View File

@ -218,12 +218,18 @@ def main() -> int:
ap.add_argument("--min-hangul-ratio", type=float, default=0.55, ap.add_argument("--min-hangul-ratio", type=float, default=0.55,
help="알파벳 문자 중 최소 한글 비율") help="알파벳 문자 중 최소 한글 비율")
ap.add_argument("--seed", type=int, default=20260819) ap.add_argument("--seed", type=int, default=20260819)
ap.add_argument("--shard-count", type=int, default=1,
help="소재 행을 서로 겹치지 않는 N개 샤드로 분할")
ap.add_argument("--shard-index", type=int, default=0,
help="이 프로세스가 담당할 0-based 샤드")
ap.add_argument("--out", type=Path, default=Path("data/training/ai_samples.jsonl")) ap.add_argument("--out", type=Path, default=Path("data/training/ai_samples.jsonl"))
ap.add_argument("--dry-run", action="store_true", ap.add_argument("--dry-run", action="store_true",
help="프롬프트만 출력하고 API 를 호출하지 않는다 (비용 0)") help="프롬프트만 출력하고 API 를 호출하지 않는다 (비용 0)")
args = ap.parse_args() args = ap.parse_args()
if args.concurrency < 1: if args.concurrency < 1:
ap.error("--concurrency는 1 이상이어야 합니다") ap.error("--concurrency는 1 이상이어야 합니다")
if args.shard_count < 1 or not 0 <= args.shard_index < args.shard_count:
ap.error("--shard-count는 1 이상, --shard-index는 0 이상 shard-count 미만이어야 합니다")
columns, rows = load_xlsx_rows(args.xlsx) columns, rows = load_xlsx_rows(args.xlsx)
if args.text_column not in columns: if args.text_column not in columns:
@ -268,7 +274,14 @@ def main() -> int:
rng = random.Random(args.seed) rng = random.Random(args.seed)
# 소재 행을 결정적으로 섞는다. 같은 seed 면 같은 순서가 나온다. # 소재 행을 결정적으로 섞는다. 같은 seed 면 같은 순서가 나온다.
order = sorted(range(len(rows)), key=lambda i: _stable_int(f"{args.seed}:{i}")) order = sorted(
(
i for i in range(len(rows))
if _stable_int(f"shard:{i}") % args.shard_count == args.shard_index
),
key=lambda i: _stable_int(f"{args.seed}:{i}"),
)
logger.info("소재 샤드 %d/%d: %d", args.shard_index, args.shard_count, len(order))
client = None client = None
if not args.dry_run: if not args.dry_run:

View File

@ -108,3 +108,9 @@ class TestDiversity:
a = gen._stable_int("abc") % len(gen.STYLE_VARIANTS) a = gen._stable_int("abc") % len(gen.STYLE_VARIANTS)
b = gen._stable_int("abc") % len(gen.STYLE_VARIANTS) b = gen._stable_int("abc") % len(gen.STYLE_VARIANTS)
assert a == b assert a == b
def test_two_shards_are_disjoint_and_complete(self):
left = {i for i in range(100) if gen._stable_int(f"shard:{i}") % 2 == 0}
right = {i for i in range(100) if gen._stable_int(f"shard:{i}") % 2 == 1}
assert not left & right
assert left | right == set(range(100))