A protein language model pre-training framework built on DeepSeek-style Llama blocks, featuring a U-Net style dual-granularity architecture with structure-aware auxiliary heads (CLE and distogram) and Fill-in-the-Middle training.
ProXYZ pre-trains protein language models that jointly learn sequence and structural signals:
- U-Net style dual-granularity model (
XYZForCausalLM) — a character-granularity transformer stack feeds into a BPE-token-granularity transformer trunk, whose output is decoded back to character granularity with U-Net skip connections. - Structure-aware auxiliary heads — on top of the character decoder, the model predicts:
- CLE (Cα-Local-Environment, 26 structure-letter alphabet) per residue
- Distogram — pairwise residue–residue distance distributions (64 bins)
- Next-token prediction — the token trunk carries the standard causal LM head.
- Fill-in-the-Middle (FIM) training — DeepSeek-Coder style SPM/PSM infilling for bidirectional context.
- Cluster-based sampling — weighted sampling (
n / (1 + log n)by cluster size) to balance sequence-cluster diversity. - HuggingFace integration — train from local files (
line/fasta/pdb) or HuggingFace Hub datasets. - Separate loss tracking — monitor standard vs. FIM loss, and each auxiliary loss, independently.
XYZForCausalLM is a three-stage U-Net over a shared embed_tokens and shared RoPE:
- Char encoder — small transformer stack at amino-acid (character) granularity.
- Token trunk — the main transformer stack at BPE-token granularity (32 layers by default). Encoder output is gathered at representative character positions (
repr_char_idx), projected back, and added to the token embeddings. - Char decoder — character-granularity stack receiving the trunk output via
scatter_addplus a U-Net skip connection from the encoder.
Three prediction heads sit on this backbone:
| Head | Input | Output | Role |
|---|---|---|---|
lm_head |
trunk output (B, T, H) | next-token logits (B, T, vocab) | causal LM objective |
cle_lm_head |
char decoder output (B, L, H_char) | CLE logits (B, L, 26) | per-residue local structure |
distogram_head |
char decoder output (B, L, H_char) | distance logits (B, L, L, 64) | pairwise residue distances |
The total loss is the next-token CE plus the auxiliary CLE / distogram CEs.
Note on "char level": the character path is a deterministic BPE tokenization (BPE dropout = 0), not per-character tokenization. The processor runs the tokenizer twice — once with BPE dropout (training regularization) and once without (for alignment between the two granularities).
install_env.sh builds a reproducible conda environment with pinned CUDA / PyTorch / flash-attn versions (Python 3.11, CUDA 12.8, PyTorch 2.11, flash-attn 2.8.3):
git clone https://github.com/bigict/ProXYZ.git
cd ProXYZ
# Create the conda environment (default name: "xyz")
bash install_env.sh
# Or specify a custom environment name
bash install_env.sh -n abc
# Activate it
conda activate xyzThe script installs the full stack: torch (cu128), flash-attn, torch_geometric, transformers, datasets, accelerate, tokenizers, graphein, biotite, biglist, lmdb, and posix-ipc. Run bash install_env.sh -h to see all options.
pip install torch transformers datasets click biotite biopython
# Optional: flash attention for best performance
pip install flash-attn --no-build-isolationBasic training with local data (standard Llama backbone):
bash train.sh protein_seqs.txtFASTA input:
bash train.sh seqs.fasta --data_format fastaTrain the U-Net model with all auxiliary heads:
bash train.sh seqs.fasta --data_format fasta \
--model_has_cle_lm_head \
--model_has_distogram_lm_head \
--model_has_char_lm_head \
--model_char_hidden_size 768 \
--model_char_intermediate_size 2064 \
--model_char_num_hidden_layers 2 \
--model_char_num_attention_heads 6Train from a HuggingFace dataset:
PYTHONPATH=src python src/proxyz/train.py \
--dataset_name your-dataset/name \
--dataset_split train \
--text_column sequence \
--tokenizer_file uniref90_30000.jsonWith --data_format pdb, CLE labels and distogram labels are derived from
experimental structures: per-protein preprocessed tensors are loaded from
$DATA_PATH/processed/<id>.pt, CLE letters are encoded from backbone /
Cβ coordinates (biotite i3d alphabet), and distogram targets are
pseudo-β (Cβ, Cα for Gly) pairwise distances binned into 64 bins over
~2.3–21.7 Å.
DATA_PATH=/path/to/pdb_data bash train.sh pdb_list.csv \
--data_format pdb \
--model_has_cle_lm_head \
--model_has_distogram_lm_head# 50% FIM, 50% standard training
bash train.sh protein_seqs.txt --fim_rate 0.5
# 100% FIM (DeepSeek-Coder style), SPM/PSM mixed
bash train.sh protein_seqs.txt --fim_rate 1.0 --fim_spm_rate 0.5Formats:
- SPM:
<BOS><fim_suffix><suffix><fim_prefix><prefix><fim_middle><middle><EOS> - PSM:
<BOS><fim_prefix><prefix><fim_suffix><suffix><fim_middle><middle><EOS>
bash train.sh protein_seqs.txt --cluster_files clusters.txtCluster file format (two columns: cluster_id, data_row_id). Sampling weight
per cluster: n / (1 + log n) where n is cluster size.
bash train.sh train.txt \
--eval_files val.txt \
--eval_strategy steps \
--eval_steps 500 \
--resume_from_checkpoint--resume_from_checkpoint restores model weights, optimizer state, LR
scheduler, and the training step from the latest checkpoint in --output_dir.
bash generate.sh # 10 seqs x 100 tokens
bash generate.sh --num_sequences 50 --num_tokens 512
bash generate.sh --prompt MVSKGE --temperature 0.8 # seeded generation
bash generate.sh --force_length --num_tokens 256 # exact length, ignore [EOS]Generated sequences are written as a timestamped FASTA under --output_dir
(default ./generated_sequences). Tokens containing X (unknown residue)
are suppressed.
Fold generated FASTA files with ESMFold to inspect structural quality:
PYTHONPATH=src python src/proxyz/evaluate.py generated_*.fasta \
--output_dir ./esmfold_pdbs--model_hidden_size(2048),--model_intermediate_size(5632),--model_num_hidden_layers(24),--model_num_attention_heads(16),--model_num_key_value_heads(4, GQA)--model_char_hidden_size(768),--model_char_intermediate_size(2064),--model_char_num_hidden_layers(2),--model_char_num_attention_heads(6)--model_has_char_lm_head/--model_has_cle_lm_head/--model_has_distogram_lm_head: enable auxiliary heads--model_use_char_position_ids: gather token positions from char positions viarepr_char_idx--max_position_embeddings(4096)
--data_format:line|fasta|pdb--tokenizer_file: BPE tokenizer JSON--tokenizer_bpe_dropout: stochastic BPE regularization rate--max_sequence_length: random-crop threshold for long sequences--text_column(text),--dataset_name/--dataset_config/--dataset_split/--dataset_eval_splitfor HuggingFace Hub--cluster_files: cluster-based sampling files
--fim_rate(0.0),--fim_spm_rate(0.5),--fim_sft_style(loss only after<fim_middle>)
--learning_rate(3e-4),--weight_decay(0.1),--warmup_steps(0)--num_train_epochs(3.0),--max_steps(-1),--per_device_train_batch_size(4),--gradient_accumulation_steps(8)
--output_dir,--logging_steps(10),--save_steps(500)--report_to(swanlab, tensorboard),--run_name,--logging_dir--resume_from_checkpoint
--attn_implementation:flash_attention_2|sdpa|eager--dataloader_num_workers(4)
ProXYZ/
├── src/proxyz/
│ ├── train.py # Training entry point
│ ├── generate.py # Sequence generation
│ ├── evaluate.py # ESMFold evaluation of generated sequences
│ ├── classify.py # Token classification
│ ├── models/
│ │ ├── configuration_xyz.py # XYZConfig
│ │ ├── modeling_xyz.py # XYZForCausalLM (U-Net + heads)
│ │ ├── modular_xyz.py # modular source of truth
│ │ └── processing_xyz.py # processor implementation
│ ├── data/ # dataset iterators, PDB feature extraction, sampler
│ └── utils/
├── script_utils/ # tokenizer / id-mapping / FIM utilities
├── assets/ # architecture diagram (svg / html)
├── install_env.sh # conda environment setup script
├── (train|generate|classify).sh # wrapper scripts
└── README.md
- Python 3.9+
- PyTorch 2.0+
- transformers 4.40+
- datasets, click
- biotite, biopython (PDB / structure features)
- flash-attn (optional, recommended)
MIT License — see LICENSE for details.
@software{proxyz2026,
title = {ProXYZ: Protein Language Model Pre-training Framework},
author = {bigict},
year = {2026},
url = {https://github.com/bigict/ProXYZ}
}- DeepSeek-Coder for FIM training methodology
- HuggingFace Transformers for the training framework