[NLP] 2-1. src 모듈화 구조 정리

Paper2Code·2025년 8월 7일

NLP Project

목록 보기
11/12
post-thumbnail

디렉토리 설계 및 최종 구현 구조

.
├── config/
├── data/
├── notebooks/
│   └── KJB/
├── script/ # gradio를 활용한 결과 분석 script 등
├── src/
│   ├── dataset/           # 데이터 전처리, 로더, 토크나이저
│   │   └── preprocess.py
│   │   └── loader.py
│
│   ├── model/             # 모델 아키텍처 및 로딩
│   │   └── base_model.py
│   │   └── lora_wrapper.py
│   │   └── peft_loader.py
│
│   ├── train/             # 학습 로직
│   │   └── train_qlora.py
│   │   └── trainer.py
│
│   ├── evaluation/        # 평가 및 메트릭
│   │   └── evaluator.py
│   │   └── metrics.py
│
│   ├── inference/         # 추론 및 생성
│   │   └── infer.py
│   │   └── generator.py
│
│   ├── util/              # 공통 유틸 함수
│   │   └── collator.py
│   │   └── logger.py
│   │   └── config_loader.py
│
│   └── main.py            # 실행 진입점 (예: argparse 기반 전체 파이프라인)
├── .env.template
├── .gitignore
├── README.md
├── requirements.txt

전체 실행 파이프라인 및 참조 파일

1. config.py
   └─ config 불러오기 (절대경로 적용)

2. train.py
   ├─ bart.py         → 모델, 토크나이저 로드
   ├─ preprocess.py   → CSV 파일 전처리
   ├─ datamodule.py   → 데이터셋 준비
   │   └─ dataset_bart.py → Dataset 클래스 정의
   ├─ seq2seqarg.py   → 학습 인자 설정
   ├─ rouge.py        → 평가 메트릭 함수
   └─ wandb.py        → 실험 로깅 초기화
   └─ huggingface Trainer → 학습 실행

3. inference.py
   ├─ bart.py         → 체크포인트에서 모델 로드
   ├─ preprocess.py   → test.csv 전처리
   ├─ datamodule.py    → test dataset 구성
   └─ test inference → 요약 생성 및 결과 저장

[1] 학습 파이프라인 (train.py)

train.py는 다음과 같은 모듈들을 호출하며 학습을 수행함:

호출 대상용도
config.py설정 파일(config.yaml)을 불러오고 절대경로 적용
bart.pyBART 모델 및 토크나이저 로딩
preprocess.py입력값 전처리 (BART에 맞게 BOS/EOS 붙이기 등)
datamodule.py학습/검증용 데이터셋 로딩 및 토크나이징
dataset_bart.pydatamodule.py에서 사용하는 custom Dataset 클래스 정의
seq2seqarg.pyhuggingface용 Seq2SeqTrainingArguments 정의
rouge.pyTrainer에서 사용하는 평가 metric 정의
wandb.pywandb 실험 로깅을 위한 초기화

[2] 추론 파이프라인 (inference.py)

inference.py는 다음처럼 구성됨:

호출 대상용도
config.pyconfig 불러오기
bart.py학습된 체크포인트에서 모델 로드
preprocess.py테스트 데이터 전처리
datamodule.py테스트셋 로딩
dataset_bart.pyinference용 dataset 사용

📦 각 파일별 설명 (의존 모듈 포함)

1. train.py
  • 학습 전반을 수행하는 스크립트
  • 내부에서 load_config, load_tokenizer_and_model_for_train, prepare_train_dataset 등 호출
  • Seq2SeqTrainer 생성 후 trainer.train() 수행
  • wandb 로깅, early stopping 포함
  • 최종 체크포인트를 config에 저장
2. inference.py
  • 테스트 데이터셋을 받아 학습된 모델로 요약문 생성
  • generate()로 예측 수행
  • 결과를 .csv로 저장
3. config.py
  • config/config.yaml을 읽고, 경로 기반 설정값 수정
  • save_config()로 수정된 설정 저장
4. `bart.py
  • 모델 및 토크나이저를 Huggingface에서 로드
  • special token 적용 → resize_token_embeddings()
5. datamodule.py
  • 전처리된 데이터를 받아서 토크나이징 및 Dataset 생성
  • train, val, test 모두 지원
  • 내부적으로 dataset_bart.py의 클래스를 사용함
6. dataset_bart.py
  • 학습/검증용: DatasetForTrain, DatasetForVal
  • 추론용: DatasetForInference
  • huggingface Trainer와 연동될 수 있는 포맷 제공
7. preprocess.py
  • CSV 파일을 받아 필요한 컬럼 추출
  • BART 입력 형태에 맞도록 bos_token, eos_token 추가
8. rouge.py
  • huggingface trainer에서 compute_metrics에 전달되어 사용
  • prediction/label을 decode → 특수 토큰 제거 → rouge.get_scores() 적용
9. seq2seqarg.py
  • Seq2SeqTrainingArguments 객체 생성
  • output_dir, logging_steps, learning_rate 등 설정값 반영
10. wandb.py
  • .env에서 WANDB_API_KEY 로드
  • 실험 이름을 모델명_타임스탬프 형식으로 생성 후 wandb에 로깅 시작
profile
As I Imagine | 이론 정리 사이트 Tistory 링크 참조 ↓ 홈 아이콘 클릭)

0개의 댓글