feat: GPU 생성 작업 샤딩 지원
This commit is contained in:
parent
78e9f8648a
commit
d0b5876ae5
@ -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:
|
||||||
|
|||||||
@ -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))
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user