Downloads · 30 days
0
niuqimeng/AGR
AGR is a machine learning model from niuqimeng. Use it for the machine learning task on the model card, and read the license before you ship it in a product. It is set up for transformers.
This model is a fine-tuned version of None. It has been trained using TRL.
Downloads · 30 days
0
Access
Public
Updated Aug 29, 2026
Repo size
16.2 GB
Likes
0
Public
Click a slice to open those files.
.safetensors16.1 GB · 99%
From the Hugging Face model README
This model is a fine-tuned version of None. It has been trained using TRL.
This is a group music recommendation model with memory-retrieval augmentation: given a group's listening history and each member's personal music preferences, it recommends the top 10 most suitable artists from a candidate list.
The model is fine-tuned from meta-llama/Meta-Llama-3-8B-Instruct — first with SFT, then with the GRPO reinforcement learning algorithm (LoRA, rlhf_type=grpo) via ms-swift. During training, each input is augmented with group memory retrieval (member listening history + common group artists).
Reward functions used in training (external plugin memory/plungin.py):
| Reward function | Weight | Description |
|---|---|---|
format_reward | 0.2 | Checks that the <think> <memory> <reasoning> <rec> tags are complete, properly closed, and in the correct order |
recommendation_reward | 0.5 | Computes Hit@k / NDCG@k between the recommendation list and the ground truth |
reasoning_reward | 0.3 | Calls the DeepSeek API to score the quality of the reasoning process |
The full workflow consists of 5 parts: Environment setup → Data preparation → Memory-enhanced preprocessing → GRPO training → Model inference.
The training scripts target a Linux server (original paths under /root/autodl-tmp) and require a GPU with at least ~16 GB VRAM (Llama-3-8B in fp16).
# Install ms-swift (training framework) and dependencies
pip install 'ms-swift[llm]' -U
pip install pandas datasets
# Create the project directory and copy the code
mkdir -p /root/autodl-tmp/grpo_memory
# Copy the memory/ folder from this repo (config.py, preprocess_dataset.py,
# memory_retriever.py, plungin.py, prompt.txt, set_paths.sh, run_grpo_enhanced.sh, csv/)
Important: edit BASE_PATH in memory/config.py and memory/set_paths.sh — all other paths are computed from it automatically:
# config.py
BASE_PATH = "/root/autodl-tmp" # <- change to your server path
Training data lives in train_data/ as a JSON array; each sample has 4 fields:
[
{
"prompt": "Based on the group's listening history and individual music preferences, analyze group and user preference features ... and recommend the top 10 suitable artists from the candidate list with ranking and reasons.",
"input": "Group info Group ID:group_common0000 Group size:2 Common artist count:7 Current common artists:Rihanna(pop,rnb,dance),MariahCarey(rnb,pop,female vocalists)...Member listening preferences Member 503:BritneySpears(tags:pop,dance,female vocalists,preference strength:9.6/10)...Candidate artist list 1.HironosukeSatou(random,unknown) 2.TheKillers(indie,indie rock,rock)...",
"output": "...reference answer...",
"ground_truth": "BritneySpears"
}
]
prompt: task instructioninput: contains the group ID, member listening preferences, and the candidate artist list; the memory retriever extracts group_xxx and member IDs from hereoutput / ground_truth: reference answer and the ground-truth artist(s), used by the reward functions to compute hit ratesGroup memory data lives in memory/csv/, three CSV files:
| File | Columns | Purpose |
|---|---|---|
user_artists.csv | userID, artistID, weight | user-artist listening relations (weight = listening weight) |
artists.csv | id, name, url | artist ID → name mapping |
tags.csv | userID, artistID, tagID, day, month, year | tagging records |
Run memory/preprocess_dataset.py. It will:
group_common0000) and member IDs from each sample's input;memory_retriever.py to retrieve each member's top-3 listening history and the group's top-3 common artists;【记忆检索结果】 (memory retrieval result) block to the original input to build the enhanced dataset;group_id, has_memory, etc.) and print enhancement statistics.cd /root/autodl-tmp/grpo_memory
python preprocess_dataset.py
# Writes the enhanced dataset to the ENHANCED_DATASET path configured in config.py
The enhanced input looks like this (original input + memory block):
...original input content...
【记忆检索结果】
【User listening history】
User 503: BritneySpears, GleeCast, Keane
User 145: JimSturgess, EllieGoulding, DavidArchuleta
【Group common artists】
Rihanna (average weight: 9)
Please make music recommendations based on the memory information above.
The training entry point is memory/run_grpo_enhanced.sh, which runs: load paths → check files → preprocess the dataset → validate JSON → start GRPO training → write a completion marker.
cd /root/autodl-tmp/grpo_memory
source set_paths.sh
bash run_grpo_enhanced.sh
The core training command (equivalent to step 5 inside the script; can also be run standalone):
swift rlhf \
--external_plugins /root/autodl-tmp/grpo_memory/plungin.py \
--reward_funcs format_reward recommendation_reward reasoning_reward \
--reward_weights 0.2 0.5 0.3 \
--rlhf_type 'grpo' \
--torch_dtype 'float16' \
--learning_rate '5e-6' \
--beta '0.001' \
--temperature 0.7 \
--top_p 0.9 \
--log_completions true \
--lora_rank 8 \
--lora_alpha 32 \
--target_modules all-linear \
--num_train_epochs '1.0' \
--per_device_train_batch_size '1' \
--gradient_accumulation_steps '8' \
--num_generations '4' \
--max_completion_length '6000' \
--overlong_filter True \
--max_length '4000' \
--save_steps '100' \
--model $SFT_MODEL \
--model_type 'llama3' \
--template 'llama3' \
--dataset $ENHANCED_DATASET \
--output_dir $GRPO_OUTPUT \
--system $PROMPT_FILE
Key parameters:
--model $SFT_MODEL: the SFT base model (SFT-tuned Llama-3-8B-Instruct)--external_plugins: loads the custom reward functions (EnhancedFormatRewardFunction / EnhancedRecommendationRewardFunction / EnhancedReasoningRewardFunction)--num_generations 4: samples 4 completions per prompt for GRPO group-wise comparison--reward_weights 0.2 0.5 0.3: weights of the three reward functions; recommendation hit rate has the highest weight--system $PROMPT_FILE: system prompt used during training (memory/prompt.txt, requires the model to output the four tags <think> <memory> <reasoning> <rec>)After training, the model is saved to $GRPO_OUTPUT (i.e. the model/ folder in this directory).
The model/ folder contains the full merged model weights (~15 GB, 5 safetensors shards) and can be loaded locally with transformers — no Hugging Face download needed.
Option A: local loading with transformers (recommended)
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
MODEL_DIR = "./model" # the model/ folder in this directory
model = AutoModelForCausalLM.from_pretrained(
MODEL_DIR,
torch_dtype=torch.bfloat16,
device_map="auto", # automatic VRAM placement; shards across multiple GPUs if available
)
tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR)
# System prompt — identical to the one used during training
system = open("./memory/prompt.txt", encoding="utf-8").read()
# Input format matches the training data 'input' field: group info + member preferences + candidate list
question = "Group info Group ID:group_common0000 Group size:2 Common artist count:7 Current common artists:Rihanna(pop,rnb,dance),MariahCarey(rnb,pop,female vocalists),KrisAllen(american idol,male vocalists,singer-songwriter) Member listening preferences Member 503:BritneySpears(tags:pop,dance,female vocalists,preference strength:9.6/10) GleeCast(tags:glee,cover,pop,preference strength:7.3/10) Keane(tags:indie,alternative,britpop,preference strength:7.0/10) Member 145:JimSturgess(tags:pop,rock,the beatles,preference strength:7.4/10) EllieGoulding(tags:electronic,female vocalists,indie,preference strength:7.2/10) DavidArchuleta(tags:pop,american idol,male vocalists,preference strength:7.2/10) Candidate artist list 1.HironosukeSatou(random,unknown) 2.TheKillers(indie,indie rock,rock) 3.KylieMinogue(pop,dance,electronic) 4.EllieGoulding(electronic,female vocalists,indie) 5.FoxesInFiction(ambient,shoegaze,dream pop) 6.Keane(indie,alternative,britpop) 7.BritneySpears(pop,dance,female vocalists) 8.Madonna(pop,dance,female vocalists) 9.ChristinaAguilera(pop,female vocalists,dance) 10.JimSturgess(pop,rock,the beatles) 11.Cryo(ebm,industrial,swedish) 12.BrandonFlowers(alternative rock,rock,indie) 13.ChrisGarneau(piano,alternative,singer-songwriter) 14.Raven(random,unknown) 15.HilaryDuff(pop,dance,female vocalists)"
messages = [
{"role": "system", "content": system},
{"role": "user", "content": question},
]
inputs = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt"
).to(model.device)
# Generation parameters match the training config (args.json)
outputs = model.generate(
inputs,
max_new_tokens=6000,
do_sample=True,
temperature=0.7,
top_p=0.9,
)
answer = tokenizer.decode(outputs[0][inputs.shape[1]:], skip_special_tokens=True)
print(answer)
# The output contains the four tags <think> <memory> <reasoning> <rec> with the full recommendation
Option B: one-liner with pipeline
from transformers import pipeline
question = "Group info Group ID:group_common0000 ... Candidate artist list 1....15...." # same format as above
generator = pipeline("text-generation", model="./model", device="cuda")
output = generator(
[{"role": "user", "content": question}],
max_new_tokens=6000,
temperature=0.7,
top_p=0.9,
return_full_text=False,
)[0]
print(output["generated_text"])
Option C: deploy with Ollama
This directory ships with an auto-generated model/Modelfile:
cd model
ollama create group-rec -f Modelfile
ollama run group-rec
Optional: enable memory augmentation at inference time as well. Training data was enhanced with memory retrieval, so applying the same enhancement to the input at inference time with memory/memory_retriever.py is recommended for best results:
import sys
sys.path.append("./memory")
from memory_retriever import GroupMemoryRetriever
retriever = GroupMemoryRetriever("./memory/csv")
enhanced_question = retriever.enhance_input(question) # appends the 【记忆检索结果】 block
# pass enhanced_question as the user content to the model
swift rlhf)